Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 1 addition & 2 deletions api/repositories/data_source/credential_repository.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
"""Actor-aware and trusted-source SQLAlchemy repository for datasource credentials."""

import json
from collections.abc import Mapping
from typing import cast

Expand All @@ -22,7 +21,7 @@
def _source_mapping(value: object) -> Mapping[str, object]:
try:
if isinstance(value, str):
return _CREDENTIALS_ADAPTER.validate_python(json.loads(value))
return _CREDENTIALS_ADAPTER.validate_json(value)
return _CREDENTIALS_ADAPTER.validate_python(value)
except (TypeError, ValueError, ValidationError):
return {}
Expand Down
3 changes: 1 addition & 2 deletions api/repositories/knowledge/document_repository.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
"""SQLAlchemy repository for tenant-owned document state."""

import json
from collections.abc import Mapping

from pydantic import TypeAdapter, ValidationError
Expand All @@ -26,7 +25,7 @@ def _mapping(value: object) -> Mapping[str, object] | None:
return None
try:
if isinstance(value, str):
return _MAPPING_ADAPTER.validate_python(json.loads(value))
return _MAPPING_ADAPTER.validate_json(value)
return _MAPPING_ADAPTER.validate_python(value)
except (TypeError, ValueError, ValidationError):
return {}
Expand Down
8 changes: 5 additions & 3 deletions api/repositories/upload_file_delivery_repository.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,8 @@
"""SQLAlchemy query adapter for public UploadFile delivery endpoints."""

import json
from typing import cast, override
from typing import override

from pydantic import TypeAdapter
from sqlalchemy import Row, Select, select
from sqlalchemy.orm import Session, sessionmaker

Expand All @@ -14,6 +14,8 @@
UploadFileDeliveryRecord,
)

_TENANT_CUSTOM_CONFIG_ADAPTER = TypeAdapter(TenantCustomConfigDict)


class UploadFileDeliveryQueryRepository(UploadFileDeliveryQuery):
def __init__(self, *, session_factory: sessionmaker[Session]) -> None:
Expand All @@ -36,7 +38,7 @@ def get_workspace_logo(self, *, workspace_id: str) -> UploadFileDeliveryRecord |
raise UploadFileDeliveryNotFoundError

custom_config = (
cast(TenantCustomConfigDict, json.loads(workspace_row.custom_config))
_TENANT_CUSTOM_CONFIG_ADAPTER.validate_json(workspace_row.custom_config)
if workspace_row.custom_config
else {}
)
Expand Down
4 changes: 1 addition & 3 deletions api/tasks/enterprise_telemetry_task.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,6 @@
dispatches them to the EnterpriseMetricHandler.
"""

import json
import logging

from celery import shared_task
Expand All @@ -31,8 +30,7 @@ def process_enterprise_telemetry(envelope_json: str) -> None:
"""
try:
# Deserialize envelope
envelope_dict = json.loads(envelope_json)
envelope = TelemetryEnvelope.model_validate(envelope_dict)
envelope = TelemetryEnvelope.model_validate_json(envelope_json)

# Process through handler
handler = EnterpriseMetricHandler()
Expand Down
16 changes: 12 additions & 4 deletions api/tasks/mail_human_input_delivery_task.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,10 @@
import json
import logging
import time
from dataclasses import dataclass
from typing import Any

import click
from celery import shared_task
from pydantic import BaseModel, ConfigDict
from sqlalchemy import select
from sqlalchemy.orm import Session, sessionmaker

Expand Down Expand Up @@ -48,14 +47,23 @@ def _build_form_link(token: str) -> str:
return f"{base_url.rstrip('/')}/form/{token}"


class _RecipientPayload(BaseModel):
"""Shape of HumanInputFormRecipient.recipient_payload written by the recipients layer."""

model_config = ConfigDict(extra="ignore")

email: str
TYPE: RecipientType


def _parse_recipient_payload(payload: str) -> tuple[str | None, RecipientType | None]:
try:
payload_dict: dict[str, Any] = json.loads(payload)
recipient = _RecipientPayload.model_validate_json(payload)
except Exception:
logger.exception("Failed to parse recipient payload")
return None, None

return payload_dict.get("email"), payload_dict.get("TYPE")
return recipient.email, recipient.TYPE


def _load_email_jobs(session: Session, form: HumanInputForm) -> list[_EmailDeliveryJob]:
Expand Down
10 changes: 5 additions & 5 deletions api/tasks/process_tenant_plugin_autoupgrade_check_task.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@

import click
from celery import shared_task
from pydantic import TypeAdapter

from core.helper.marketplace import batch_fetch_plugin_manifests, get_plugin_pkg_url
from core.plugin.entities.marketplace import MarketplacePluginDeclaration, MarketplacePluginSnapshot
Expand All @@ -25,6 +26,9 @@
CACHE_REDIS_KEY_PREFIX = "plugin_autoupgrade_check_task:cached_plugin_snapshot:"
CACHE_REDIS_TTL = 60 * 60 # 1 hour

# A cached miss is stored as JSON null, so None has to stay a valid payload.
_CACHED_MANIFEST_ADAPTER = TypeAdapter(typing.Union[MarketplacePluginSnapshot, None])


def _get_redis_cache_key(plugin_id: str) -> str:
"""Generate Redis cache key for plugin manifest."""
Expand All @@ -45,11 +49,7 @@ def _get_cached_manifest(plugin_id: str) -> typing.Union[MarketplacePluginSnapsh
if cached_data is None:
return False

cached_json = json.loads(cached_data)
if cached_json is None:
return None

return MarketplacePluginSnapshot.model_validate(cached_json)
return _CACHED_MANIFEST_ADAPTER.validate_json(cached_data)
except Exception:
logger.exception("Failed to get cached manifest for plugin %s", plugin_id)
return False
Expand Down
7 changes: 5 additions & 2 deletions api/tasks/rag_pipeline/priority_rag_pipeline_run_task.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,4 @@
import contextvars
import json
import logging
import time
import uuid
Expand All @@ -10,6 +9,7 @@
import click
from celery import shared_task # type: ignore
from flask import current_app, g
from pydantic import TypeAdapter
from sqlalchemy import select
from sqlalchemy.orm import Session, sessionmaker

Expand All @@ -31,6 +31,9 @@

logger = logging.getLogger(__name__)

# Each entry is validated into RagPipelineInvokeEntity inside the worker.
_INVOKE_ENTITIES_ADAPTER = TypeAdapter(list[dict[str, Any]])


@shared_task(queue="priority_pipeline")
def priority_rag_pipeline_run_task(
Expand All @@ -50,7 +53,7 @@ def priority_rag_pipeline_run_task(
rag_pipeline_invoke_entities_content = FileService(db.engine).get_file_content(
rag_pipeline_invoke_entities_file_id
)
rag_pipeline_invoke_entities = json.loads(rag_pipeline_invoke_entities_content)
rag_pipeline_invoke_entities = _INVOKE_ENTITIES_ADAPTER.validate_json(rag_pipeline_invoke_entities_content)

logger.info("tenant %s received %d rag pipeline invoke entities", tenant_id, len(rag_pipeline_invoke_entities))

Expand Down
7 changes: 5 additions & 2 deletions api/tasks/rag_pipeline/rag_pipeline_run_task.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,4 @@
import contextvars
import json
import logging
import time
import uuid
Expand All @@ -11,6 +10,7 @@
import click
from celery import group, shared_task
from flask import current_app, g
from pydantic import TypeAdapter
from sqlalchemy import select
from sqlalchemy.orm import sessionmaker

Expand All @@ -32,6 +32,9 @@

logger = logging.getLogger(__name__)

# Each entry is validated into RagPipelineInvokeEntity inside the worker.
_INVOKE_ENTITIES_ADAPTER = TypeAdapter(list[dict[str, Any]])


def chunked(iterable: Sequence, size: int):
it = iter(iterable)
Expand All @@ -56,7 +59,7 @@ def rag_pipeline_run_task(
rag_pipeline_invoke_entities_content = FileService(db.engine).get_file_content(
rag_pipeline_invoke_entities_file_id
)
rag_pipeline_invoke_entities = json.loads(rag_pipeline_invoke_entities_content)
rag_pipeline_invoke_entities = _INVOKE_ENTITIES_ADAPTER.validate_json(rag_pipeline_invoke_entities_content)

logger.info("tenant %s received %d rag pipeline invoke entities", tenant_id, len(rag_pipeline_invoke_entities))

Expand Down
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import json
from collections.abc import Sequence
from datetime import UTC, datetime, timedelta
from types import SimpleNamespace
Expand All @@ -7,7 +8,7 @@
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session, sessionmaker

from models.human_input import HumanInputForm
from models.human_input import HumanInputForm, RecipientType
from tasks import mail_human_input_delivery_task as task_module


Expand Down Expand Up @@ -178,3 +179,23 @@ def test_dispatch_human_input_email_task_sanitizes_subject(
)

assert mail.sent[0]["subject"] == "Notice BCC:attacker@example.com Alert"


def test_parse_recipient_payload_reads_email_and_type():
payload = json.dumps({"email": "user@example.com", "TYPE": RecipientType.EMAIL_MEMBER.value})
Comment thread
asukaminato0721 marked this conversation as resolved.

assert task_module._parse_recipient_payload(payload) == ("user@example.com", RecipientType.EMAIL_MEMBER)


@pytest.mark.parametrize(
"payload",
[
"not json",
json.dumps({"TYPE": RecipientType.EMAIL_MEMBER.value}),
json.dumps({"email": "user@example.com", "TYPE": "unknown-channel"}),
json.dumps({"email": 42, "TYPE": RecipientType.EMAIL_MEMBER.value}),
Comment thread
asukaminato0721 marked this conversation as resolved.
],
)
def test_parse_recipient_payload_rejects_unusable_payloads(payload: str):
"""Payloads the recipients layer cannot produce must degrade to a skipped recipient."""
assert task_module._parse_recipient_payload(payload) == (None, None)
Loading