From 6676481af21676eeca7881cd64d3a8e7eb7b7d5a Mon Sep 17 00:00:00 2001 From: suguanYang Date: Mon, 13 Jul 2026 15:52:58 +0800 Subject: [PATCH 01/19] Expand OSS telemetry with ADR-0004 schema, heartbeats, and usage aggregates. Formalize the anonymous self-hosted emitter taxonomy so operators get reliable adoption and health signals without collecting document or user content. Co-authored-by: Cursor --- README.md | 18 + apps/api/.env.example | 12 + apps/api/main.py | 37 +- .../test_self_hosted_telemetry_contract.py | 187 +++++++- .../0004-anonymous-self-hosted-telemetry.md | 110 +++++ docs/adr/README.md | 10 + .../shared/services/telemetry/__init__.py | 11 +- .../shared/services/telemetry/aggregates.py | 398 +++++++++++++++--- .../shared/services/telemetry/config.py | 4 +- .../shared/services/telemetry/events.py | 151 ++++++- .../shared/services/telemetry/runtime.py | 185 +++++++- 11 files changed, 1059 insertions(+), 64 deletions(-) create mode 100644 docs/adr/0004-anonymous-self-hosted-telemetry.md diff --git a/README.md b/README.md index 8681c6028..442a540f5 100644 --- a/README.md +++ b/README.md @@ -244,6 +244,24 @@ make check - External dependency guide: [docs/external-services.md](docs/external-services.md) +- Architecture decisions: + [docs/adr/README.md](docs/adr/README.md) + +## Telemetry + +Self-hosted Knowhere emits **anonymous** product telemetry to PostHog so Ontos +operators can understand OSS adoption (install liveness, usage aggregates, +client/document mix). Events never include filenames, prompts, emails, IPs, or +geo. Schema and allowlists are locked in +[ADR-0004](docs/adr/0004-anonymous-self-hosted-telemetry.md). + +Telemetry is **default-on**. To opt out, set: + +```bash +TELEMETRY_ENABLED=false +``` + +Related settings live in `apps/api/.env.example` under `TELEMETRY_*`. ## Citation diff --git a/apps/api/.env.example b/apps/api/.env.example index f4dc6ce8d..04f297b40 100644 --- a/apps/api/.env.example +++ b/apps/api/.env.example @@ -30,6 +30,18 @@ TMP_PATH=/tmp/knowhere # Optional or development-only: observability and local dashboard wiring LOGFIRE_TOKEN= +# Anonymous self-hosted product telemetry (PostHog). Default-on; opt out with false. +# See docs/adr/0004-anonymous-self-hosted-telemetry.md and the README Telemetry section. +TELEMETRY_ENABLED=true +# TELEMETRY_POSTHOG_HOST=https://us.i.posthog.com +# TELEMETRY_POSTHOG_PROJECT_KEY= +# TELEMETRY_INSTALLATION_ID= +# TELEMETRY_INSTALLATION_ID_PATH=/data/secrets/telemetry-installation-id +# TELEMETRY_BATCH_SIZE=20 +# TELEMETRY_REQUEST_TIMEOUT_SECONDS=2.0 +# TELEMETRY_DEPLOYMENT_MODE=self_hosted +# TELEMETRY_AGGREGATE_INTERVAL_SECONDS=300 + # Required for local startup: database DATABASE_URL=postgresql+asyncpg://root:root123@localhost:5432/Knowhere DB_SSL_MODE=disable diff --git a/apps/api/main.py b/apps/api/main.py index 7140fab0d..3e5475661 100644 --- a/apps/api/main.py +++ b/apps/api/main.py @@ -70,10 +70,23 @@ async def lifespan(app: FastAPI): await load_rules(session) logger.info("rate limit rules loaded at startup; restart the pod to apply changes") + import time + from shared.services.telemetry.aggregates import ( start_self_hosted_aggregate_telemetry, ) - from shared.services.telemetry.runtime import start_self_hosted_telemetry + from shared.services.telemetry.runtime import ( + build_postgres_health_probe, + build_redis_health_probe, + start_self_hosted_heartbeat_telemetry, + start_self_hosted_telemetry, + ) + + telemetry_started_at = time.monotonic() + + async def _redis_ping() -> bool: + redis_service = redis_pool_manager.get_redis_service() + return await redis_service.ping() telemetry_runtime = await start_self_hosted_telemetry( settings, @@ -86,6 +99,7 @@ async def lifespan(app: FastAPI): app.state.self_hosted_telemetry_client = None app.state.self_hosted_telemetry_config = None app.state.self_hosted_aggregate_telemetry_runner = None + app.state.self_hosted_heartbeat_telemetry_runner = None else: telemetry_client, telemetry_config = telemetry_runtime app.state.self_hosted_telemetry_client = telemetry_client @@ -99,6 +113,16 @@ async def lifespan(app: FastAPI): api_metrics=app.state.self_hosted_api_telemetry_metrics, ) ) + app.state.self_hosted_heartbeat_telemetry_runner = ( + await start_self_hosted_heartbeat_telemetry( + settings, + telemetry_client=telemetry_client, + config=telemetry_config, + started_at_monotonic=telemetry_started_at, + postgres_probe=build_postgres_health_probe(get_db_context), + redis_probe=build_redis_health_probe(_redis_ping), + ) + ) mcp_server = getattr(app.state, "retrieval_mcp_server", None) mcp_session_manager = getattr(mcp_server, "session_manager", None) @@ -114,13 +138,20 @@ async def lifespan(app: FastAPI): from shared.services.telemetry.aggregates import ( stop_self_hosted_aggregate_telemetry, ) - from shared.services.telemetry.runtime import stop_self_hosted_telemetry + from shared.services.telemetry.runtime import ( + stop_self_hosted_heartbeat_telemetry, + stop_self_hosted_telemetry, + ) + await stop_self_hosted_heartbeat_telemetry( + getattr(app.state, "self_hosted_heartbeat_telemetry_runner", None) + ) await stop_self_hosted_aggregate_telemetry( getattr(app.state, "self_hosted_aggregate_telemetry_runner", None) ) await stop_self_hosted_telemetry( - getattr(app.state, "self_hosted_telemetry_client", None) + getattr(app.state, "self_hosted_telemetry_client", None), + config=getattr(app.state, "self_hosted_telemetry_config", None), ) except Exception as e: logger.error(f"self-hosted telemetry shutdown failed: {e}") diff --git a/apps/api/tests/contract/test_self_hosted_telemetry_contract.py b/apps/api/tests/contract/test_self_hosted_telemetry_contract.py index 54152c258..b974e7613 100644 --- a/apps/api/tests/contract/test_self_hosted_telemetry_contract.py +++ b/apps/api/tests/contract/test_self_hosted_telemetry_contract.py @@ -26,11 +26,22 @@ stop_self_hosted_aggregate_telemetry, ) from shared.services.telemetry.events import ( + SCHEMA_VERSION, + build_base_event_properties, + compute_success_rate, + count_to_bucket, get_allowed_telemetry_event_names, + normalize_client_name, + normalize_document_type, + normalize_source_type, sanitize_event_properties, + uptime_seconds_to_bucket, ) from shared.services.telemetry.identity import get_or_create_installation_id - +from shared.services.telemetry.runtime import ( + SelfHostedHeartbeatTelemetryRunner, + stop_self_hosted_telemetry, +) def test_installation_id_is_generated_once(tmp_path: Path) -> None: installation_id_path = tmp_path / "telemetry-installation-id" @@ -99,9 +110,175 @@ def test_aggregate_event_names_are_allowed() -> None: "self_hosted_worker_aggregate", "self_hosted_api_aggregate", "self_hosted_provider_aggregate", + "self_hosted_document_type_aggregate", + "self_hosted_client_aggregate", }.issubset(get_allowed_telemetry_event_names()) +def test_normalize_document_type_and_client_name() -> None: + assert normalize_document_type("report.PDF") == "pdf" + assert normalize_document_type("photo.jpeg") == "image" + assert normalize_document_type("notes.htm") == "html" + assert normalize_document_type("secret.xyz") == "other" + assert normalize_document_type(None) == "other" + assert normalize_client_name("node-sdk") == "node-sdk" + assert normalize_client_name("CLI") == "cli" + assert normalize_client_name("custom-bot") == "other" + assert normalize_source_type("direct_upload") == "file" + assert normalize_source_type("url") == "url" + assert normalize_source_type("demo") == "other" + + +def test_success_rate_and_count_buckets() -> None: + assert compute_success_rate(9, 1) == 0.9 + assert compute_success_rate(0, 0) == 0.0 + assert count_to_bucket(0) == "0" + assert count_to_bucket(7) == "1-10" + assert count_to_bucket(50) == "11-100" + assert count_to_bucket(101) == "100+" + assert uptime_seconds_to_bucket(30) == "0m-5m" + assert uptime_seconds_to_bucket(3600) == "1h-24h" + + +def test_usage_and_document_type_properties_strip_sensitive_values() -> None: + usage_properties = sanitize_event_properties( + "self_hosted_usage_aggregate", + { + "app_version": "1.2.3", + "window_seconds": 86_400, + "success_rate_24h": 0.9, + "source_file_jobs_24h": 2, + "email": "user@example.com", + "document_name": "private.pdf", + "source_file_name": "private.pdf", + }, + ) + document_type_properties = sanitize_event_properties( + "self_hosted_document_type_aggregate", + { + "document_type": "pdf", + "jobs_created_24h": 1, + "source_file_name": "private.pdf", + "email": "user@example.com", + }, + ) + client_properties = sanitize_event_properties( + "self_hosted_client_aggregate", + { + "created_by_client": "cli", + "jobs_created_24h": 1, + "client_version": "9.9.9", + "email": "user@example.com", + }, + ) + + assert usage_properties == { + "app_version": "1.2.3", + "window_seconds": 86_400, + "success_rate_24h": 0.9, + "source_file_jobs_24h": 2, + } + assert document_type_properties == { + "document_type": "pdf", + "jobs_created_24h": 1, + } + assert client_properties == { + "created_by_client": "cli", + "jobs_created_24h": 1, + } + + +@pytest.mark.asyncio +async def test_heartbeat_emit_once_includes_health_and_uptime( + tmp_path: Path, +) -> None: + posthog_client = _FakePostHogClient() + config = _build_config(tmp_path) + telemetry_client = TelemetryClient(config, posthog_client=posthog_client) + await telemetry_client.start() + + async def postgres_probe() -> bool: + return True + + async def redis_probe() -> bool: + return False + + runner = SelfHostedHeartbeatTelemetryRunner( + config=config, + telemetry_client=telemetry_client, + settings=_HeartbeatSettings(), + interval_seconds=60, + started_at_monotonic=0.0, + postgres_probe=postgres_probe, + redis_probe=redis_probe, + ) + await runner.emit_once() + await telemetry_client.stop() + + assert len(posthog_client.captured_events) == 1 + captured = posthog_client.captured_events[0] + assert captured.event_name == "self_hosted_instance_heartbeat" + properties = cast(dict[str, object], captured.kwargs["properties"]) + assert properties["api_healthy"] is True + assert properties["postgres_healthy"] is True + assert properties["redis_healthy"] is False + assert properties["uptime_bucket"] in { + "0m-5m", + "5m-1h", + "1h-24h", + "24h-7d", + "7d+", + } + assert properties["schema_version"] == SCHEMA_VERSION + + +@pytest.mark.asyncio +async def test_shutdown_includes_base_event_properties(tmp_path: Path) -> None: + posthog_client = _FakePostHogClient() + config = _build_config(tmp_path) + telemetry_client = TelemetryClient(config, posthog_client=posthog_client) + await telemetry_client.start() + + await stop_self_hosted_telemetry(telemetry_client, config=config) + + assert len(posthog_client.captured_events) == 1 + captured = posthog_client.captured_events[0] + assert captured.event_name == "self_hosted_instance_shutdown" + assert captured.kwargs["properties"] == { + **build_base_event_properties(config), + "$process_person_profile": False, + } + + +def test_usage_aggregate_allowlist_includes_v2_keys() -> None: + properties = sanitize_event_properties( + "self_hosted_usage_aggregate", + { + "success_rate_24h": 1.0, + "job_duration_p95_seconds_24h": 12.5, + "has_webhooks_24h": True, + "has_retrieval_24h": False, + "jobs_created_bucket": "1-10", + "pages_processed_bucket": "0", + "source_file_jobs_24h": 1, + "source_url_jobs_24h": 0, + "source_other_jobs_24h": 0, + "filename": "secret.pdf", + }, + ) + assert set(properties) == { + "success_rate_24h", + "job_duration_p95_seconds_24h", + "has_webhooks_24h", + "has_retrieval_24h", + "jobs_created_bucket", + "pages_processed_bucket", + "source_file_jobs_24h", + "source_url_jobs_24h", + "source_other_jobs_24h", + } + + def test_self_hosted_telemetry_defaults_to_enabled( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -428,6 +605,13 @@ class _AggregateSettings: TELEMETRY_AGGREGATE_INTERVAL_SECONDS: int = 60 +@dataclass(frozen=True) +class _HeartbeatSettings: + API_STANDALONE_MODE_ENABLED: bool = False + BILLING_ENABLED: bool = False + TELEMETRY_AGGREGATE_INTERVAL_SECONDS: int = 60 + + class _FailingSessionContext(AbstractAsyncContextManager[AsyncSession]): async def __aenter__(self) -> AsyncSession: raise RuntimeError("database unavailable") @@ -517,4 +701,5 @@ def _build_config( environment="production", app_env="production", service_name="knowhere-api", + schema_version=SCHEMA_VERSION, ) diff --git a/docs/adr/0004-anonymous-self-hosted-telemetry.md b/docs/adr/0004-anonymous-self-hosted-telemetry.md new file mode 100644 index 000000000..bb51f2603 --- /dev/null +++ b/docs/adr/0004-anonymous-self-hosted-telemetry.md @@ -0,0 +1,110 @@ +# 0004 Anonymous Self-Hosted Telemetry + +## Status + +Accepted + +## Context + +Self-hosted Knowhere installs already emit anonymous PostHog events so Ontos +operators can understand OSS adoption. The v1 emitter covered instance +lifecycle and coarse aggregates, but lacked: + +- a locked privacy and allowlist contract +- real periodic health heartbeats +- SaaS `/usage`-parity KPIs (success rate, p95 duration, source mix) +- document-type and client-type mix without leaking filenames or free-form metadata + +Operators need a stable schema and metric catalog before the metrics dashboard +and official clients depend on new events. + +## Decision + +### Purpose and opt-out + +Anonymous self-hosted telemetry exists for Ontos operators measuring OSS / +self-hosted adoption. It is **default-on**. Operators opt out with +`TELEMETRY_ENABLED=false`. Transport remains PostHog; Logfire/OTEL are out of +scope. + +### Privacy bounds + +Events must never include filenames, prompts, emails, IPs, geo, document +content, or arbitrary customer metadata keys. Only allowlisted scalar property +names may leave the box. Free-form `client_version` is not emitted in aggregate +events (cardinality); app version stays on base properties only. + +### Schema version + +`schema_version = 2026-07-telemetry-v2` + +### Event catalog + +Keep the existing eight event names and add two: + +| Event | Role | +| --- | --- | +| `self_hosted_instance_started` | Install boot | +| `self_hosted_instance_heartbeat` | Periodic liveness + health | +| `self_hosted_instance_shutdown` | Graceful stop (includes base props) | +| `self_hosted_usage_aggregate` | Fleet usage snapshot (24h window) | +| `self_hosted_retrieval_aggregate` | Retrieval activity | +| `self_hosted_worker_aggregate` | Worker backlog / completion | +| `self_hosted_api_aggregate` | In-process API request counters | +| `self_hosted_provider_aggregate` | Parse/retrieval provider activity | +| `self_hosted_document_type_aggregate` | Per allowlisted document type | +| `self_hosted_client_aggregate` | Per allowlisted created_by_client | + +### Allowlists + +- `document_type`: `pdf|docx|doc|xlsx|xls|pptx|ppt|csv|txt|md|html|image|other` + - Extension from `job_metadata->>'source_file_name'` only; map + `png|jpg|jpeg|gif|webp|tiff` → `image`; everything else → `other` +- `created_by_client`: `cli|node-sdk|dashboard|notebook|mcp|api|other` + - From `job_metadata #>> '{document_metadata,created_by_client}'` +- `source_type`: `file|url|other` + +### Success rate + +`success_rate_24h = done / (done + failed)` over the same 24h window, as a +float in **0–1**. Exclude non-terminal statuses. When `(done + failed) = 0`, +emit `0.0`. + +### Metric IDs (dashboard catalog) + +| Metric ID | Source | +| --- | --- | +| `oss.active_installs_7d` / `oss.active_installs_30d` | heartbeat distinct installs | +| `oss.new_installs_*` | started | +| `oss.retention.w0_w1` / `oss.retention.w0_w4` | started ∩ heartbeat | +| `oss.usage.jobs_created_24h` | usage aggregate | +| `oss.usage.jobs_done_24h` / `oss.usage.jobs_failed_24h` | usage / worker | +| `oss.usage.success_rate_24h` | usage aggregate | +| `oss.usage.pages_processed_24h` | usage aggregate | +| `oss.usage.job_duration_avg_seconds_24h` | worker aggregate | +| `oss.usage.job_duration_p95_seconds_24h` | usage aggregate | +| `oss.usage.backlog_*` | worker pending/running/… | +| `oss.usage.source_*` | `source_{file,url,other}_jobs_24h` | +| `oss.client_*` | client aggregate | +| `oss.document_type_*` | document type aggregate | +| `oss.health.*` | heartbeat health fields | +| `oss.fleet.version_*` / `oss.fleet.flags_*` | base props last-known | + +### Capability flags and buckets + +Usage aggregate may emit: + +- `has_webhooks_24h`, `has_retrieval_24h` (bool) +- `jobs_created_bucket`, `pages_processed_bucket` as + `0|1-10|11-100|100+` + +## Consequences + +- Emitter code must bump `schema_version`, extend property allowlists, emit + real heartbeats, and add document-type / client aggregates before the + metrics-dashboard Telemetry UI can rely on them. +- Official clients should populate `document_metadata.created_by_client` + (separate decision surface); until then dashboards must tolerate + `other`-heavy client mix. +- Changing allowlists or the success-rate formula requires a new ADR or an + explicit schema_version bump. diff --git a/docs/adr/README.md b/docs/adr/README.md index 3d0cb4fd0..68e15b221 100644 --- a/docs/adr/README.md +++ b/docs/adr/README.md @@ -10,3 +10,13 @@ Use this shape: - Context - Decision - Consequences + +## Index + +| ADR | Title | +| --- | --- | +| [0001](0001-keep-routes-and-worker-tasks-as-adapters.md) | Keep routes and worker tasks as adapters | +| [0002](0002-use-typed-workflow-outcomes.md) | Use typed workflow outcomes | +| [0003](0003-keep-retrieval-workflow-policy-explicit.md) | Keep retrieval workflow policy explicit | +| [0004](0004-anonymous-self-hosted-telemetry.md) | Anonymous self-hosted telemetry | +| \ No newline at end of file diff --git a/packages/shared-python/shared/services/telemetry/__init__.py b/packages/shared-python/shared/services/telemetry/__init__.py index 047ce17d6..ee2fa6003 100644 --- a/packages/shared-python/shared/services/telemetry/__init__.py +++ b/packages/shared-python/shared/services/telemetry/__init__.py @@ -1,18 +1,27 @@ """Anonymous self-hosted telemetry helpers.""" from .client import TelemetryClient -from .config import TelemetryRuntimeConfig, build_telemetry_config +from .config import SCHEMA_VERSION, TelemetryRuntimeConfig, build_telemetry_config from .events import ( + build_base_event_properties, build_instance_event_properties, get_allowed_telemetry_event_names, + normalize_client_name, + normalize_document_type, + normalize_source_type, ) from .identity import get_or_create_installation_id __all__ = [ + "SCHEMA_VERSION", "TelemetryClient", "TelemetryRuntimeConfig", "build_telemetry_config", + "build_base_event_properties", "build_instance_event_properties", "get_allowed_telemetry_event_names", "get_or_create_installation_id", + "normalize_client_name", + "normalize_document_type", + "normalize_source_type", ] diff --git a/packages/shared-python/shared/services/telemetry/aggregates.py b/packages/shared-python/shared/services/telemetry/aggregates.py index ddb5af9b7..80634ccc5 100644 --- a/packages/shared-python/shared/services/telemetry/aggregates.py +++ b/packages/shared-python/shared/services/telemetry/aggregates.py @@ -13,7 +13,15 @@ from .api_metrics import ApiRequestTelemetryMetrics, ApiRequestMetricsSnapshot from .client import TelemetryClient from .config import TelemetryRuntimeConfig, TelemetrySettings -from .events import TelemetryProperties, build_base_event_properties +from .events import ( + TelemetryProperties, + build_base_event_properties, + compute_success_rate, + count_to_bucket, + normalize_client_name, + normalize_document_type, + normalize_source_type, +) AGGREGATE_WINDOW_SECONDS = 24 * 60 * 60 AGGREGATE_ADVISORY_LOCK_ID = 0x4B4E4F5748455245 @@ -52,14 +60,14 @@ def __init__( async def emit_once(self) -> None: """Collect and emit all aggregate snapshots once.""" - event_properties = await collect_self_hosted_aggregate_event_properties( + event_captures = await collect_self_hosted_aggregate_event_captures( config=self._config, db_session_factory=self._db_session_factory, api_metrics=self._api_metrics, window_seconds=AGGREGATE_WINDOW_SECONDS, api_window_seconds=self._interval_seconds, ) - for event_name, properties in event_properties.items(): + for event_name, properties in event_captures: self._telemetry_client.capture(event_name, properties) def start(self) -> None: @@ -126,61 +134,164 @@ async def stop_self_hosted_aggregate_telemetry( await runner.stop() -async def collect_self_hosted_aggregate_event_properties( +async def collect_self_hosted_aggregate_event_captures( *, config: TelemetryRuntimeConfig, db_session_factory: DatabaseSessionFactory, api_metrics: ApiRequestTelemetryMetrics, window_seconds: int = AGGREGATE_WINDOW_SECONDS, api_window_seconds: int = AGGREGATE_WINDOW_SECONDS, -) -> dict[str, TelemetryProperties]: - """Collect aggregate event properties without including customer content.""" - event_properties = { - "self_hosted_api_aggregate": _collect_api_aggregate( - config, - api_metrics.snapshot_and_reset(), - api_window_seconds, - ), - } +) -> list[tuple[str, TelemetryProperties]]: + """Collect aggregate captures without including customer content.""" + captures: list[tuple[str, TelemetryProperties]] = [ + ( + "self_hosted_api_aggregate", + _collect_api_aggregate( + config, + api_metrics.snapshot_and_reset(), + api_window_seconds, + ), + ) + ] async with db_session_factory() as session: lock_acquired = await _try_aggregate_advisory_lock(session) if not lock_acquired: - return event_properties + return captures try: - event_properties.update( - { - "self_hosted_usage_aggregate": await _collect_usage_aggregate( - session, - config, - window_seconds, - ), - "self_hosted_retrieval_aggregate": await _collect_retrieval_aggregate( - session, - config, - window_seconds, - ), - "self_hosted_worker_aggregate": await _collect_worker_aggregate( - session, - config, - window_seconds, - ), - "self_hosted_provider_aggregate": await _collect_provider_aggregate( + captures.append( + ( + "self_hosted_usage_aggregate", + await _collect_usage_aggregate(session, config, window_seconds), + ) + ) + captures.append( + ( + "self_hosted_retrieval_aggregate", + await _collect_retrieval_aggregate( session, config, window_seconds, ), - } + ) ) - return event_properties + captures.append( + ( + "self_hosted_worker_aggregate", + await _collect_worker_aggregate(session, config, window_seconds), + ) + ) + captures.append( + ( + "self_hosted_provider_aggregate", + await _collect_provider_aggregate(session, config, window_seconds), + ) + ) + for properties in await _collect_document_type_aggregates( + session, + config, + window_seconds, + ): + captures.append(("self_hosted_document_type_aggregate", properties)) + for properties in await _collect_client_aggregates( + session, + config, + window_seconds, + ): + captures.append(("self_hosted_client_aggregate", properties)) + return captures finally: await _release_aggregate_advisory_lock(session) +async def collect_self_hosted_aggregate_event_properties( + *, + config: TelemetryRuntimeConfig, + db_session_factory: DatabaseSessionFactory, + api_metrics: ApiRequestTelemetryMetrics, + window_seconds: int = AGGREGATE_WINDOW_SECONDS, + api_window_seconds: int = AGGREGATE_WINDOW_SECONDS, +) -> dict[str, TelemetryProperties]: + """Collect singleton aggregate event properties (compat helper for tests).""" + captures = await collect_self_hosted_aggregate_event_captures( + config=config, + db_session_factory=db_session_factory, + api_metrics=api_metrics, + window_seconds=window_seconds, + api_window_seconds=api_window_seconds, + ) + event_properties: dict[str, TelemetryProperties] = {} + for event_name, properties in captures: + # Keep first capture for singleton events; multi-row events are omitted + # from this dict helper (use collect_self_hosted_aggregate_event_captures). + if event_name in { + "self_hosted_document_type_aggregate", + "self_hosted_client_aggregate", + }: + continue + event_properties[event_name] = properties + return event_properties + + async def _collect_usage_aggregate( session: AsyncSession, config: TelemetryRuntimeConfig, window_seconds: int, ) -> TelemetryProperties: + completed_jobs_24h = await _int_scalar( + session, + _windowed_count_sql("jobs", "updated_at", "status = 'done'"), + window_seconds, + ) + failed_jobs_24h = await _int_scalar( + session, + _windowed_count_sql("jobs", "updated_at", "status = 'failed'"), + window_seconds, + ) + jobs_created_24h = await _int_scalar( + session, + _windowed_count_sql("jobs", "created_at"), + window_seconds, + ) + pages_processed_24h = await _int_scalar( + session, + f""" + SELECT COALESCE(SUM(page_count), 0) + FROM jobs + WHERE updated_at >= {_window_start_expression()} + AND status = 'done' + """, + window_seconds, + ) + source_counts = await _collect_source_type_counts(session, window_seconds) + has_webhooks_24h = ( + await _int_scalar( + session, + f""" + SELECT CASE + WHEN EXISTS ( + SELECT 1 FROM jobs + WHERE created_at >= {_window_start_expression()} + AND webhook_enabled = true + ) + OR EXISTS ( + SELECT 1 FROM webhook_logs + WHERE created_at >= {_window_start_expression()} + ) + THEN 1 ELSE 0 + END + """, + window_seconds, + ) + > 0 + ) + has_retrieval_24h = ( + await _int_scalar( + session, + _windowed_count_sql("retrieval_runs", "created_at"), + window_seconds, + ) + > 0 + ) properties = _base_aggregate_properties(config, window_seconds) properties.update( { @@ -190,23 +301,31 @@ async def _collect_usage_aggregate( "SELECT COUNT(*) FROM api_keys WHERE is_active = true", ), "total_jobs": await _int_scalar(session, "SELECT COUNT(*) FROM jobs"), - "jobs_created_24h": await _int_scalar( - session, - _windowed_count_sql("jobs", "created_at"), - window_seconds, - ), + "jobs_created_24h": jobs_created_24h, + "jobs_created_bucket": count_to_bucket(jobs_created_24h), "active_jobs": await _int_scalar( session, "SELECT COUNT(*) FROM jobs WHERE status IN ('waiting-file', 'pending', 'running', 'converting')", ), - "completed_jobs_24h": await _int_scalar( - session, - _windowed_count_sql("jobs", "updated_at", "status = 'done'"), - window_seconds, + "completed_jobs_24h": completed_jobs_24h, + "failed_jobs_24h": failed_jobs_24h, + "success_rate_24h": compute_success_rate( + completed_jobs_24h, + failed_jobs_24h, ), - "failed_jobs_24h": await _int_scalar( + "job_duration_p95_seconds_24h": await _float_scalar( session, - _windowed_count_sql("jobs", "updated_at", "status = 'failed'"), + f""" + SELECT COALESCE( + percentile_cont(0.95) WITHIN GROUP ( + ORDER BY EXTRACT(EPOCH FROM (updated_at - created_at)) + ), + 0 + ) + FROM jobs + WHERE updated_at >= {_window_start_expression()} + AND status IN ('done', 'failed') + """, window_seconds, ), "total_documents": await _int_scalar( @@ -225,16 +344,8 @@ async def _collect_usage_aggregate( session, "SELECT COUNT(*) FROM job_chunks", ), - "pages_processed_24h": await _int_scalar( - session, - f""" - SELECT COALESCE(SUM(page_count), 0) - FROM jobs - WHERE updated_at >= {_window_start_expression()} - AND status = 'done' - """, - window_seconds, - ), + "pages_processed_24h": pages_processed_24h, + "pages_processed_bucket": count_to_bucket(pages_processed_24h), "credits_charged_24h": await _int_scalar( session, f""" @@ -244,11 +355,188 @@ async def _collect_usage_aggregate( """, window_seconds, ), + "has_webhooks_24h": has_webhooks_24h, + "has_retrieval_24h": has_retrieval_24h, + "source_file_jobs_24h": source_counts["file"], + "source_url_jobs_24h": source_counts["url"], + "source_other_jobs_24h": source_counts["other"], } ) return properties +async def _collect_source_type_counts( + session: AsyncSession, + window_seconds: int, +) -> dict[str, int]: + counts = {"file": 0, "url": 0, "other": 0} + result = await session.execute( + text( + f""" + SELECT source_type, COUNT(*) AS job_count + FROM jobs + WHERE created_at >= {_window_start_expression()} + GROUP BY source_type + """ + ), + {"window_seconds": window_seconds}, + ) + for row in result.mappings(): + source_type = normalize_source_type(cast(str | None, row["source_type"])) + counts[source_type] = counts.get(source_type, 0) + int(row["job_count"] or 0) + return counts + + +async def _collect_document_type_aggregates( + session: AsyncSession, + config: TelemetryRuntimeConfig, + window_seconds: int, +) -> list[TelemetryProperties]: + """Collect one aggregate per allowlisted document type with activity.""" + result = await session.execute( + text( + f""" + SELECT + job_metadata->>'source_file_name' AS source_file_name, + COUNT(*) FILTER ( + WHERE created_at >= {_window_start_expression()} + ) AS jobs_created_24h, + COUNT(*) FILTER ( + WHERE updated_at >= {_window_start_expression()} + AND status = 'done' + ) AS jobs_done_24h, + COUNT(*) FILTER ( + WHERE updated_at >= {_window_start_expression()} + AND status = 'failed' + ) AS jobs_failed_24h, + COALESCE( + SUM(page_count) FILTER ( + WHERE updated_at >= {_window_start_expression()} + AND status = 'done' + ), + 0 + ) AS pages_processed_24h + FROM jobs + WHERE created_at >= {_window_start_expression()} + OR ( + updated_at >= {_window_start_expression()} + AND status IN ('done', 'failed') + ) + GROUP BY job_metadata->>'source_file_name' + """ + ), + {"window_seconds": window_seconds}, + ) + merged: dict[str, dict[str, int]] = {} + for row in result.mappings(): + document_type = normalize_document_type( + cast(str | None, row["source_file_name"]) + ) + bucket = merged.setdefault( + document_type, + { + "jobs_created_24h": 0, + "jobs_done_24h": 0, + "jobs_failed_24h": 0, + "pages_processed_24h": 0, + }, + ) + bucket["jobs_created_24h"] += int(row["jobs_created_24h"] or 0) + bucket["jobs_done_24h"] += int(row["jobs_done_24h"] or 0) + bucket["jobs_failed_24h"] += int(row["jobs_failed_24h"] or 0) + bucket["pages_processed_24h"] += int(row["pages_processed_24h"] or 0) + + properties_list: list[TelemetryProperties] = [] + for document_type, counts in sorted(merged.items()): + if not any(counts.values()): + continue + properties = _base_aggregate_properties(config, window_seconds) + properties.update( + { + "document_type": document_type, + "jobs_created_24h": counts["jobs_created_24h"], + "jobs_done_24h": counts["jobs_done_24h"], + "jobs_failed_24h": counts["jobs_failed_24h"], + "pages_processed_24h": counts["pages_processed_24h"], + "success_rate_24h": compute_success_rate( + counts["jobs_done_24h"], + counts["jobs_failed_24h"], + ), + } + ) + properties_list.append(properties) + return properties_list + + +async def _collect_client_aggregates( + session: AsyncSession, + config: TelemetryRuntimeConfig, + window_seconds: int, +) -> list[TelemetryProperties]: + """Collect one aggregate per allowlisted created_by_client with activity.""" + result = await session.execute( + text( + f""" + SELECT + job_metadata #>> '{{document_metadata,created_by_client}}' AS created_by_client, + COUNT(*) FILTER ( + WHERE created_at >= {_window_start_expression()} + ) AS jobs_created_24h, + COUNT(*) FILTER ( + WHERE updated_at >= {_window_start_expression()} + AND status = 'done' + ) AS jobs_done_24h, + COUNT(*) FILTER ( + WHERE updated_at >= {_window_start_expression()} + AND status = 'failed' + ) AS jobs_failed_24h + FROM jobs + WHERE created_at >= {_window_start_expression()} + OR ( + updated_at >= {_window_start_expression()} + AND status IN ('done', 'failed') + ) + GROUP BY job_metadata #>> '{{document_metadata,created_by_client}}' + """ + ), + {"window_seconds": window_seconds}, + ) + merged: dict[str, dict[str, int]] = {} + for row in result.mappings(): + client_name = normalize_client_name(cast(str | None, row["created_by_client"])) + bucket = merged.setdefault( + client_name, + { + "jobs_created_24h": 0, + "jobs_done_24h": 0, + "jobs_failed_24h": 0, + }, + ) + bucket["jobs_created_24h"] += int(row["jobs_created_24h"] or 0) + bucket["jobs_done_24h"] += int(row["jobs_done_24h"] or 0) + bucket["jobs_failed_24h"] += int(row["jobs_failed_24h"] or 0) + + properties_list: list[TelemetryProperties] = [] + for client_name, counts in sorted(merged.items()): + if not any(counts.values()): + continue + properties = _base_aggregate_properties(config, window_seconds) + properties.update( + { + "created_by_client": client_name, + "jobs_created_24h": counts["jobs_created_24h"], + "jobs_done_24h": counts["jobs_done_24h"], + "jobs_failed_24h": counts["jobs_failed_24h"], + "success_rate_24h": compute_success_rate( + counts["jobs_done_24h"], + counts["jobs_failed_24h"], + ), + } + ) + properties_list.append(properties) + return properties_list + + async def _collect_retrieval_aggregate( session: AsyncSession, config: TelemetryRuntimeConfig, diff --git a/packages/shared-python/shared/services/telemetry/config.py b/packages/shared-python/shared/services/telemetry/config.py index 328f27060..d318f6921 100644 --- a/packages/shared-python/shared/services/telemetry/config.py +++ b/packages/shared-python/shared/services/telemetry/config.py @@ -6,6 +6,8 @@ from pathlib import Path from typing import Protocol +SCHEMA_VERSION = "2026-07-telemetry-v2" + class TelemetrySettings(Protocol): """Subset of app settings required by anonymous telemetry.""" @@ -41,7 +43,7 @@ class TelemetryRuntimeConfig: environment: str app_env: str service_name: str - schema_version: str = "2026-06-telemetry-v1" + schema_version: str = SCHEMA_VERSION @property def is_ready(self) -> bool: diff --git a/packages/shared-python/shared/services/telemetry/events.py b/packages/shared-python/shared/services/telemetry/events.py index 287e8ab98..ff3ad8b90 100644 --- a/packages/shared-python/shared/services/telemetry/events.py +++ b/packages/shared-python/shared/services/telemetry/events.py @@ -4,13 +4,55 @@ import os from collections.abc import Mapping +from pathlib import PurePosixPath from typing import TypeAlias, cast -from .config import TelemetryRuntimeConfig +from .config import SCHEMA_VERSION, TelemetryRuntimeConfig TelemetryPropertyValue: TypeAlias = str | int | float | bool | None TelemetryProperties: TypeAlias = dict[str, TelemetryPropertyValue] +DOCUMENT_TYPES = frozenset( + { + "pdf", + "docx", + "doc", + "xlsx", + "xls", + "pptx", + "ppt", + "csv", + "txt", + "md", + "html", + "image", + "other", + } +) + +CLIENT_NAMES = frozenset( + { + "cli", + "node-sdk", + "dashboard", + "notebook", + "mcp", + "api", + "other", + } +) + +SOURCE_TYPES = frozenset({"file", "url", "other"}) + +_IMAGE_EXTENSIONS = frozenset({"png", "jpg", "jpeg", "gif", "webp", "tiff"}) +_HTML_EXTENSIONS = frozenset({"html", "htm"}) + +_COUNT_BUCKETS = ( + (0, "0"), + (10, "1-10"), + (100, "11-100"), +) + _BASE_PROPERTY_NAMES = frozenset( { "app_env", @@ -66,8 +108,17 @@ "completed_jobs_24h", "credits_charged_24h", "failed_jobs_24h", + "has_retrieval_24h", + "has_webhooks_24h", + "job_duration_p95_seconds_24h", "jobs_created_24h", + "jobs_created_bucket", "pages_processed_24h", + "pages_processed_bucket", + "source_file_jobs_24h", + "source_other_jobs_24h", + "source_url_jobs_24h", + "success_rate_24h", "total_document_chunks", "total_documents", "total_job_chunks", @@ -128,6 +179,27 @@ "webhook_delivery_failures_24h", } ), + "self_hosted_document_type_aggregate": _AGGREGATE_PROPERTY_NAMES + | frozenset( + { + "document_type", + "jobs_created_24h", + "jobs_done_24h", + "jobs_failed_24h", + "pages_processed_24h", + "success_rate_24h", + } + ), + "self_hosted_client_aggregate": _AGGREGATE_PROPERTY_NAMES + | frozenset( + { + "created_by_client", + "jobs_created_24h", + "jobs_done_24h", + "jobs_failed_24h", + "success_rate_24h", + } + ), } @@ -136,6 +208,83 @@ def get_allowed_telemetry_event_names() -> frozenset[str]: return frozenset(_EVENT_PROPERTY_NAMES.keys()) +def normalize_document_type(extension_or_filename: str | None) -> str: + """Map a filename or extension to an allowlisted document_type.""" + if not extension_or_filename: + return "other" + raw = extension_or_filename.strip().lower() + if not raw: + return "other" + # Accept either "pdf" or "report.pdf" / "path/report.PDF". + if "/" in raw or "\\" in raw or "." in raw: + suffix = PurePosixPath(raw.replace("\\", "/")).suffix + extension = suffix.lstrip(".") + else: + extension = raw.lstrip(".") + if not extension: + return "other" + if extension in _IMAGE_EXTENSIONS: + return "image" + if extension in _HTML_EXTENSIONS: + return "html" + if extension in DOCUMENT_TYPES: + return extension + return "other" + + +def normalize_client_name(raw: str | None) -> str: + """Map a created_by_client value to an allowlisted client name.""" + if raw is None: + return "other" + normalized = raw.strip().lower() + if normalized in CLIENT_NAMES: + return normalized + return "other" + + +def normalize_source_type(raw: str | None) -> str: + """Map a job source_type value to an allowlisted source_type.""" + if raw is None: + return "other" + normalized = raw.strip().lower() + if normalized == "direct_upload": + return "file" + if normalized in SOURCE_TYPES: + return normalized + return "other" + + +def compute_success_rate(done: int, failed: int) -> float: + """Return done / (done + failed) as a float in 0–1.""" + terminal = max(done, 0) + max(failed, 0) + if terminal == 0: + return 0.0 + return max(done, 0) / terminal + + +def count_to_bucket(count: int) -> str: + """Bucket a non-negative count into 0|1-10|11-100|100+.""" + value = max(count, 0) + for upper_bound, label in _COUNT_BUCKETS: + if value <= upper_bound: + return label + return "100+" + + +def uptime_seconds_to_bucket(uptime_seconds: float) -> str: + """Bucket process uptime for heartbeat events.""" + seconds = max(uptime_seconds, 0.0) + if seconds < 5 * 60: + return "0m-5m" + if seconds < 60 * 60: + return "5m-1h" + if seconds < 24 * 60 * 60: + return "1h-24h" + if seconds < 7 * 24 * 60 * 60: + return "24h-7d" + return "7d+" + + def build_instance_event_properties( config: TelemetryRuntimeConfig, *, diff --git a/packages/shared-python/shared/services/telemetry/runtime.py b/packages/shared-python/shared/services/telemetry/runtime.py index d1bba0ed8..184970cef 100644 --- a/packages/shared-python/shared/services/telemetry/runtime.py +++ b/packages/shared-python/shared/services/telemetry/runtime.py @@ -2,15 +2,108 @@ from __future__ import annotations +import asyncio +import time +from collections.abc import Awaitable, Callable +from contextlib import AbstractAsyncContextManager from pathlib import Path +from typing import Protocol from loguru import logger +from sqlalchemy import text +from sqlalchemy.ext.asyncio import AsyncSession from .client import TelemetryClient from .config import TelemetryRuntimeConfig, TelemetrySettings, build_telemetry_config -from .events import build_instance_event_properties +from .events import ( + build_base_event_properties, + build_instance_event_properties, + uptime_seconds_to_bucket, +) from .identity import get_or_create_installation_id +HealthProbe = Callable[[], Awaitable[bool]] + + +class DatabaseSessionFactory(Protocol): + """Factory for app-owned async database sessions.""" + + def __call__(self) -> AbstractAsyncContextManager[AsyncSession]: + """Return an async context manager yielding an AsyncSession.""" + raise NotImplementedError + + +class TelemetryHeartbeatSettings(TelemetrySettings, Protocol): + TELEMETRY_AGGREGATE_INTERVAL_SECONDS: int + + +class SelfHostedHeartbeatTelemetryRunner: + """Periodically emits real health heartbeats for anonymous telemetry.""" + + def __init__( + self, + *, + config: TelemetryRuntimeConfig, + telemetry_client: TelemetryClient, + settings: TelemetrySettings, + interval_seconds: int, + started_at_monotonic: float, + postgres_probe: HealthProbe | None = None, + redis_probe: HealthProbe | None = None, + ) -> None: + self._config = config + self._telemetry_client = telemetry_client + self._settings = settings + self._interval_seconds = interval_seconds + self._started_at_monotonic = started_at_monotonic + self._postgres_probe = postgres_probe + self._redis_probe = redis_probe + self._task: asyncio.Task[None] | None = None + + async def emit_once(self) -> None: + """Probe dependencies and capture one heartbeat event.""" + postgres_healthy = await _run_health_probe(self._postgres_probe, default=False) + redis_healthy = await _run_health_probe(self._redis_probe, default=False) + uptime_seconds = time.monotonic() - self._started_at_monotonic + self._telemetry_client.capture( + "self_hosted_instance_heartbeat", + build_instance_event_properties( + self._config, + api_standalone_mode_enabled=self._settings.API_STANDALONE_MODE_ENABLED, + billing_enabled=self._settings.BILLING_ENABLED, + api_healthy=True, + postgres_healthy=postgres_healthy, + redis_healthy=redis_healthy, + uptime_bucket=uptime_seconds_to_bucket(uptime_seconds), + ), + ) + + def start(self) -> None: + """Start the periodic heartbeat loop.""" + if self._task is not None: + return + self._task = asyncio.create_task( + self._run(), + name="self-hosted-heartbeat-telemetry", + ) + + async def stop(self) -> None: + """Stop the periodic heartbeat loop.""" + task = self._task + if task is None: + return + task.cancel() + await asyncio.gather(task, return_exceptions=True) + self._task = None + + async def _run(self) -> None: + while True: + await asyncio.sleep(self._interval_seconds) + try: + await self.emit_once() + except Exception as exc: + logger.warning(f"anonymous heartbeat telemetry failed: {exc}") + async def start_self_hosted_telemetry( settings: TelemetrySettings, @@ -70,11 +163,99 @@ async def start_self_hosted_telemetry( return telemetry_client, config +async def start_self_hosted_heartbeat_telemetry( + settings: TelemetryHeartbeatSettings, + *, + telemetry_client: TelemetryClient | None, + config: TelemetryRuntimeConfig | None, + started_at_monotonic: float | None = None, + postgres_probe: HealthProbe | None = None, + redis_probe: HealthProbe | None = None, +) -> SelfHostedHeartbeatTelemetryRunner | None: + """Start periodic real heartbeats when the base self-hosted client is active.""" + if telemetry_client is None or config is None: + return None + interval_seconds = max(settings.TELEMETRY_AGGREGATE_INTERVAL_SECONDS, 60) + runner = SelfHostedHeartbeatTelemetryRunner( + config=config, + telemetry_client=telemetry_client, + settings=settings, + interval_seconds=interval_seconds, + started_at_monotonic=( + time.monotonic() if started_at_monotonic is None else started_at_monotonic + ), + postgres_probe=postgres_probe, + redis_probe=redis_probe, + ) + try: + runner.start() + except Exception as exc: + logger.warning(f"anonymous heartbeat telemetry start failed: {exc}") + return None + logger.info("anonymous self-hosted heartbeat telemetry scheduled") + return runner + + +async def stop_self_hosted_heartbeat_telemetry( + runner: SelfHostedHeartbeatTelemetryRunner | None, +) -> None: + """Stop heartbeat telemetry if it was started.""" + if runner is None: + return + await runner.stop() + + async def stop_self_hosted_telemetry( telemetry_client: TelemetryClient | None, + *, + config: TelemetryRuntimeConfig | None = None, ) -> None: """Flush and stop anonymous self-hosted telemetry.""" if telemetry_client is None: return - telemetry_client.capture("self_hosted_instance_shutdown", {}) + shutdown_properties = ( + build_base_event_properties(config) if config is not None else {} + ) + telemetry_client.capture("self_hosted_instance_shutdown", shutdown_properties) await telemetry_client.stop() + + +def build_postgres_health_probe( + db_session_factory: DatabaseSessionFactory, +) -> HealthProbe: + """Return a probe that runs SELECT 1 against Postgres.""" + + async def _probe() -> bool: + try: + async with db_session_factory() as session: + result = await session.execute(text("SELECT 1")) + return result.scalar_one_or_none() == 1 + except Exception: + return False + + return _probe + + +def build_redis_health_probe(redis_ping: Callable[[], Awaitable[bool]]) -> HealthProbe: + """Return a probe that pings Redis through the provided callable.""" + + async def _probe() -> bool: + try: + return bool(await redis_ping()) + except Exception: + return False + + return _probe + + +async def _run_health_probe( + probe: HealthProbe | None, + *, + default: bool, +) -> bool: + if probe is None: + return default + try: + return bool(await probe()) + except Exception: + return False From 4361e45936e2bda92c315a6673ec6173ffacdc9c Mon Sep 17 00:00:00 2001 From: suguanYang Date: Mon, 13 Jul 2026 16:11:32 +0800 Subject: [PATCH 02/19] Potential fix for pull request finding 'CodeQL / Unused import' Co-authored-by: Copilot Autofix powered by AI <62310815+github-advanced-security[bot]@users.noreply.github.com> --- packages/shared-python/shared/services/telemetry/events.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/packages/shared-python/shared/services/telemetry/events.py b/packages/shared-python/shared/services/telemetry/events.py index ff3ad8b90..44ce071cb 100644 --- a/packages/shared-python/shared/services/telemetry/events.py +++ b/packages/shared-python/shared/services/telemetry/events.py @@ -7,7 +7,7 @@ from pathlib import PurePosixPath from typing import TypeAlias, cast -from .config import SCHEMA_VERSION, TelemetryRuntimeConfig +from .config import TelemetryRuntimeConfig TelemetryPropertyValue: TypeAlias = str | int | float | bool | None TelemetryProperties: TypeAlias = dict[str, TelemetryPropertyValue] From b4a9403f418181a748089b5d98f8163f3b315f7d Mon Sep 17 00:00:00 2001 From: suguanYang Date: Mon, 13 Jul 2026 16:26:52 +0800 Subject: [PATCH 03/19] Rename telemetry events from self_hosted_* to oss_*. Co-authored-by: Cursor --- .../test_self_hosted_telemetry_contract.py | 60 +++++++++---------- .../0004-anonymous-self-hosted-telemetry.md | 22 +++---- .../shared/services/telemetry/aggregates.py | 18 +++--- .../shared/services/telemetry/events.py | 20 +++---- .../shared/services/telemetry/runtime.py | 8 +-- 5 files changed, 64 insertions(+), 64 deletions(-) diff --git a/apps/api/tests/contract/test_self_hosted_telemetry_contract.py b/apps/api/tests/contract/test_self_hosted_telemetry_contract.py index b974e7613..a58401f19 100644 --- a/apps/api/tests/contract/test_self_hosted_telemetry_contract.py +++ b/apps/api/tests/contract/test_self_hosted_telemetry_contract.py @@ -86,7 +86,7 @@ def test_explicit_installation_id_must_be_uuid(tmp_path: Path) -> None: def test_telemetry_properties_strip_unknown_and_non_scalar_values() -> None: properties = sanitize_event_properties( - "self_hosted_instance_heartbeat", + "oss_instance_heartbeat", { "app_version": "1.2.3", "api_healthy": True, @@ -105,13 +105,13 @@ def test_telemetry_properties_strip_unknown_and_non_scalar_values() -> None: def test_aggregate_event_names_are_allowed() -> None: assert { - "self_hosted_usage_aggregate", - "self_hosted_retrieval_aggregate", - "self_hosted_worker_aggregate", - "self_hosted_api_aggregate", - "self_hosted_provider_aggregate", - "self_hosted_document_type_aggregate", - "self_hosted_client_aggregate", + "oss_usage_aggregate", + "oss_retrieval_aggregate", + "oss_worker_aggregate", + "oss_api_aggregate", + "oss_provider_aggregate", + "oss_document_type_aggregate", + "oss_client_aggregate", }.issubset(get_allowed_telemetry_event_names()) @@ -142,7 +142,7 @@ def test_success_rate_and_count_buckets() -> None: def test_usage_and_document_type_properties_strip_sensitive_values() -> None: usage_properties = sanitize_event_properties( - "self_hosted_usage_aggregate", + "oss_usage_aggregate", { "app_version": "1.2.3", "window_seconds": 86_400, @@ -154,7 +154,7 @@ def test_usage_and_document_type_properties_strip_sensitive_values() -> None: }, ) document_type_properties = sanitize_event_properties( - "self_hosted_document_type_aggregate", + "oss_document_type_aggregate", { "document_type": "pdf", "jobs_created_24h": 1, @@ -163,7 +163,7 @@ def test_usage_and_document_type_properties_strip_sensitive_values() -> None: }, ) client_properties = sanitize_event_properties( - "self_hosted_client_aggregate", + "oss_client_aggregate", { "created_by_client": "cli", "jobs_created_24h": 1, @@ -217,7 +217,7 @@ async def redis_probe() -> bool: assert len(posthog_client.captured_events) == 1 captured = posthog_client.captured_events[0] - assert captured.event_name == "self_hosted_instance_heartbeat" + assert captured.event_name == "oss_instance_heartbeat" properties = cast(dict[str, object], captured.kwargs["properties"]) assert properties["api_healthy"] is True assert properties["postgres_healthy"] is True @@ -243,7 +243,7 @@ async def test_shutdown_includes_base_event_properties(tmp_path: Path) -> None: assert len(posthog_client.captured_events) == 1 captured = posthog_client.captured_events[0] - assert captured.event_name == "self_hosted_instance_shutdown" + assert captured.event_name == "oss_instance_shutdown" assert captured.kwargs["properties"] == { **build_base_event_properties(config), "$process_person_profile": False, @@ -252,7 +252,7 @@ async def test_shutdown_includes_base_event_properties(tmp_path: Path) -> None: def test_usage_aggregate_allowlist_includes_v2_keys() -> None: properties = sanitize_event_properties( - "self_hosted_usage_aggregate", + "oss_usage_aggregate", { "success_rate_24h": 1.0, "job_duration_p95_seconds_24h": 12.5, @@ -309,7 +309,7 @@ def test_self_hosted_telemetry_env_can_disable_and_override_key( def test_aggregate_properties_strip_sensitive_values() -> None: properties = sanitize_event_properties( - "self_hosted_usage_aggregate", + "oss_usage_aggregate", { "app_version": "1.2.3", "window_seconds": 86_400, @@ -361,8 +361,8 @@ async def test_api_aggregate_uses_interval_window_when_global_lock_unavailable( api_window_seconds=300, ) - assert set(properties) == {"self_hosted_api_aggregate"} - api_properties = properties["self_hosted_api_aggregate"] + assert set(properties) == {"oss_api_aggregate"} + api_properties = properties["oss_api_aggregate"] assert api_properties["window_seconds"] == 300 assert api_properties["api_requests_total"] == 1 assert api_properties["api_requests_2xx"] == 1 @@ -375,7 +375,7 @@ def test_telemetry_client_filters_posthog_sdk_properties_after_capture( sanitized_message = telemetry_client._sanitize_posthog_message( { - "event": "self_hosted_api_aggregate", + "event": "oss_api_aggregate", "properties": { "app_version": "1.2.3", "api_requests_total": 1, @@ -420,7 +420,7 @@ async def test_telemetry_client_sends_anonymous_posthog_capture( await telemetry_client.start() queued = telemetry_client.capture( - "self_hosted_instance_heartbeat", + "oss_instance_heartbeat", { "app_version": "1.2.3", "api_healthy": True, @@ -432,7 +432,7 @@ async def test_telemetry_client_sends_anonymous_posthog_capture( assert queued is True assert len(posthog_client.captured_events) == 1 captured_event = posthog_client.captured_events[0] - assert captured_event.event_name == "self_hosted_instance_heartbeat" + assert captured_event.event_name == "oss_instance_heartbeat" assert captured_event.kwargs["distinct_id"] == ( "550e8400-e29b-41d4-a716-446655440000" ) @@ -455,7 +455,7 @@ async def test_telemetry_client_sends_aggregate_events(tmp_path: Path) -> None: await telemetry_client.start() queued = telemetry_client.capture( - "self_hosted_api_aggregate", + "oss_api_aggregate", { "app_version": "1.2.3", "window_seconds": 86_400, @@ -469,7 +469,7 @@ async def test_telemetry_client_sends_aggregate_events(tmp_path: Path) -> None: assert queued is True assert len(posthog_client.captured_events) == 1 captured_event = posthog_client.captured_events[0] - assert captured_event.event_name == "self_hosted_api_aggregate" + assert captured_event.event_name == "oss_api_aggregate" assert captured_event.kwargs["properties"] == { "app_version": "1.2.3", "window_seconds": 86_400, @@ -516,7 +516,7 @@ async def test_telemetry_client_respects_batch_size(tmp_path: Path) -> None: await telemetry_client.start() for index in range(3): telemetry_client.capture( - "self_hosted_instance_heartbeat", + "oss_instance_heartbeat", { "app_version": f"1.2.{index}", }, @@ -524,9 +524,9 @@ async def test_telemetry_client_respects_batch_size(tmp_path: Path) -> None: await telemetry_client.stop() assert [event.event_name for event in posthog_client.captured_events] == [ - "self_hosted_instance_heartbeat", - "self_hosted_instance_heartbeat", - "self_hosted_instance_heartbeat", + "oss_instance_heartbeat", + "oss_instance_heartbeat", + "oss_instance_heartbeat", ] assert posthog_client.flush_count == 1 @@ -542,7 +542,7 @@ async def test_telemetry_client_flush_before_start_does_not_deadlock( ) telemetry_client.capture( - "self_hosted_instance_heartbeat", + "oss_instance_heartbeat", { "app_version": "1.2.3", }, @@ -550,7 +550,7 @@ async def test_telemetry_client_flush_before_start_does_not_deadlock( await telemetry_client.flush() await telemetry_client.start() telemetry_client.capture( - "self_hosted_instance_heartbeat", + "oss_instance_heartbeat", { "app_version": "1.2.4", }, @@ -570,7 +570,7 @@ async def test_telemetry_client_does_not_restart_after_stop(tmp_path: Path) -> N await telemetry_client.start() telemetry_client.capture( - "self_hosted_instance_heartbeat", + "oss_instance_heartbeat", { "app_version": "1.2.3", }, @@ -578,7 +578,7 @@ async def test_telemetry_client_does_not_restart_after_stop(tmp_path: Path) -> N await telemetry_client.stop() await telemetry_client.start() queued_after_stop = telemetry_client.capture( - "self_hosted_instance_heartbeat", + "oss_instance_heartbeat", { "app_version": "1.2.4", }, diff --git a/docs/adr/0004-anonymous-self-hosted-telemetry.md b/docs/adr/0004-anonymous-self-hosted-telemetry.md index bb51f2603..d13e257a9 100644 --- a/docs/adr/0004-anonymous-self-hosted-telemetry.md +++ b/docs/adr/0004-anonymous-self-hosted-telemetry.md @@ -40,20 +40,20 @@ events (cardinality); app version stays on base properties only. ### Event catalog -Keep the existing eight event names and add two: +Event names use the `oss_` prefix (not `self_hosted_`). Catalog: | Event | Role | | --- | --- | -| `self_hosted_instance_started` | Install boot | -| `self_hosted_instance_heartbeat` | Periodic liveness + health | -| `self_hosted_instance_shutdown` | Graceful stop (includes base props) | -| `self_hosted_usage_aggregate` | Fleet usage snapshot (24h window) | -| `self_hosted_retrieval_aggregate` | Retrieval activity | -| `self_hosted_worker_aggregate` | Worker backlog / completion | -| `self_hosted_api_aggregate` | In-process API request counters | -| `self_hosted_provider_aggregate` | Parse/retrieval provider activity | -| `self_hosted_document_type_aggregate` | Per allowlisted document type | -| `self_hosted_client_aggregate` | Per allowlisted created_by_client | +| `oss_instance_started` | Install boot | +| `oss_instance_heartbeat` | Periodic liveness + health | +| `oss_instance_shutdown` | Graceful stop (includes base props) | +| `oss_usage_aggregate` | Fleet usage snapshot (24h window) | +| `oss_retrieval_aggregate` | Retrieval activity | +| `oss_worker_aggregate` | Worker backlog / completion | +| `oss_api_aggregate` | In-process API request counters | +| `oss_provider_aggregate` | Parse/retrieval provider activity | +| `oss_document_type_aggregate` | Per allowlisted document type | +| `oss_client_aggregate` | Per allowlisted created_by_client | ### Allowlists diff --git a/packages/shared-python/shared/services/telemetry/aggregates.py b/packages/shared-python/shared/services/telemetry/aggregates.py index 80634ccc5..80dfe9ec7 100644 --- a/packages/shared-python/shared/services/telemetry/aggregates.py +++ b/packages/shared-python/shared/services/telemetry/aggregates.py @@ -145,7 +145,7 @@ async def collect_self_hosted_aggregate_event_captures( """Collect aggregate captures without including customer content.""" captures: list[tuple[str, TelemetryProperties]] = [ ( - "self_hosted_api_aggregate", + "oss_api_aggregate", _collect_api_aggregate( config, api_metrics.snapshot_and_reset(), @@ -160,13 +160,13 @@ async def collect_self_hosted_aggregate_event_captures( try: captures.append( ( - "self_hosted_usage_aggregate", + "oss_usage_aggregate", await _collect_usage_aggregate(session, config, window_seconds), ) ) captures.append( ( - "self_hosted_retrieval_aggregate", + "oss_retrieval_aggregate", await _collect_retrieval_aggregate( session, config, @@ -176,13 +176,13 @@ async def collect_self_hosted_aggregate_event_captures( ) captures.append( ( - "self_hosted_worker_aggregate", + "oss_worker_aggregate", await _collect_worker_aggregate(session, config, window_seconds), ) ) captures.append( ( - "self_hosted_provider_aggregate", + "oss_provider_aggregate", await _collect_provider_aggregate(session, config, window_seconds), ) ) @@ -191,13 +191,13 @@ async def collect_self_hosted_aggregate_event_captures( config, window_seconds, ): - captures.append(("self_hosted_document_type_aggregate", properties)) + captures.append(("oss_document_type_aggregate", properties)) for properties in await _collect_client_aggregates( session, config, window_seconds, ): - captures.append(("self_hosted_client_aggregate", properties)) + captures.append(("oss_client_aggregate", properties)) return captures finally: await _release_aggregate_advisory_lock(session) @@ -224,8 +224,8 @@ async def collect_self_hosted_aggregate_event_properties( # Keep first capture for singleton events; multi-row events are omitted # from this dict helper (use collect_self_hosted_aggregate_event_captures). if event_name in { - "self_hosted_document_type_aggregate", - "self_hosted_client_aggregate", + "oss_document_type_aggregate", + "oss_client_aggregate", }: continue event_properties[event_name] = properties diff --git a/packages/shared-python/shared/services/telemetry/events.py b/packages/shared-python/shared/services/telemetry/events.py index 44ce071cb..533266f76 100644 --- a/packages/shared-python/shared/services/telemetry/events.py +++ b/packages/shared-python/shared/services/telemetry/events.py @@ -89,8 +89,8 @@ ) _EVENT_PROPERTY_NAMES: dict[str, frozenset[str]] = { - "self_hosted_instance_started": frozenset(), - "self_hosted_instance_heartbeat": frozenset( + "oss_instance_started": frozenset(), + "oss_instance_heartbeat": frozenset( { "api_healthy", "postgres_healthy", @@ -98,8 +98,8 @@ "uptime_bucket", } ), - "self_hosted_instance_shutdown": frozenset(), - "self_hosted_usage_aggregate": _AGGREGATE_PROPERTY_NAMES + "oss_instance_shutdown": frozenset(), + "oss_usage_aggregate": _AGGREGATE_PROPERTY_NAMES | frozenset( { "active_api_keys", @@ -126,7 +126,7 @@ "total_users", } ), - "self_hosted_retrieval_aggregate": _AGGREGATE_PROPERTY_NAMES + "oss_retrieval_aggregate": _AGGREGATE_PROPERTY_NAMES | frozenset( { "retrieval_cache_hits_24h", @@ -140,7 +140,7 @@ "retrieval_tokens_24h", } ), - "self_hosted_worker_aggregate": _AGGREGATE_PROPERTY_NAMES + "oss_worker_aggregate": _AGGREGATE_PROPERTY_NAMES | frozenset( { "job_duration_avg_seconds_24h", @@ -152,7 +152,7 @@ "jobs_waiting_file", } ), - "self_hosted_api_aggregate": _AGGREGATE_PROPERTY_NAMES + "oss_api_aggregate": _AGGREGATE_PROPERTY_NAMES | frozenset( { "api_latency_avg_ms", @@ -164,7 +164,7 @@ "api_requests_total", } ), - "self_hosted_provider_aggregate": _AGGREGATE_PROPERTY_NAMES + "oss_provider_aggregate": _AGGREGATE_PROPERTY_NAMES | frozenset( { "parse_agent_errors_24h", @@ -179,7 +179,7 @@ "webhook_delivery_failures_24h", } ), - "self_hosted_document_type_aggregate": _AGGREGATE_PROPERTY_NAMES + "oss_document_type_aggregate": _AGGREGATE_PROPERTY_NAMES | frozenset( { "document_type", @@ -190,7 +190,7 @@ "success_rate_24h", } ), - "self_hosted_client_aggregate": _AGGREGATE_PROPERTY_NAMES + "oss_client_aggregate": _AGGREGATE_PROPERTY_NAMES | frozenset( { "created_by_client", diff --git a/packages/shared-python/shared/services/telemetry/runtime.py b/packages/shared-python/shared/services/telemetry/runtime.py index 184970cef..5a63fab4a 100644 --- a/packages/shared-python/shared/services/telemetry/runtime.py +++ b/packages/shared-python/shared/services/telemetry/runtime.py @@ -66,7 +66,7 @@ async def emit_once(self) -> None: redis_healthy = await _run_health_probe(self._redis_probe, default=False) uptime_seconds = time.monotonic() - self._started_at_monotonic self._telemetry_client.capture( - "self_hosted_instance_heartbeat", + "oss_instance_heartbeat", build_instance_event_properties( self._config, api_standalone_mode_enabled=self._settings.API_STANDALONE_MODE_ENABLED, @@ -145,9 +145,9 @@ async def start_self_hosted_telemetry( api_standalone_mode_enabled=settings.API_STANDALONE_MODE_ENABLED, billing_enabled=settings.BILLING_ENABLED, ) - telemetry_client.capture("self_hosted_instance_started", base_properties) + telemetry_client.capture("oss_instance_started", base_properties) telemetry_client.capture( - "self_hosted_instance_heartbeat", + "oss_instance_heartbeat", build_instance_event_properties( config, api_standalone_mode_enabled=settings.API_STANDALONE_MODE_ENABLED, @@ -216,7 +216,7 @@ async def stop_self_hosted_telemetry( shutdown_properties = ( build_base_event_properties(config) if config is not None else {} ) - telemetry_client.capture("self_hosted_instance_shutdown", shutdown_properties) + telemetry_client.capture("oss_instance_shutdown", shutdown_properties) await telemetry_client.stop() From d6e3f42393d737929cef64cf013b9adea18f9c84 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Mon, 13 Jul 2026 17:19:38 +0800 Subject: [PATCH 04/19] Fix telemetry contract test SCHEMA_VERSION import. Co-authored-by: Cursor --- apps/api/tests/contract/test_self_hosted_telemetry_contract.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/apps/api/tests/contract/test_self_hosted_telemetry_contract.py b/apps/api/tests/contract/test_self_hosted_telemetry_contract.py index a58401f19..69e352ae9 100644 --- a/apps/api/tests/contract/test_self_hosted_telemetry_contract.py +++ b/apps/api/tests/contract/test_self_hosted_telemetry_contract.py @@ -17,7 +17,7 @@ BaseConfig, ) from shared.services.telemetry.client import TelemetryClient -from shared.services.telemetry.config import TelemetryRuntimeConfig +from shared.services.telemetry.config import SCHEMA_VERSION, TelemetryRuntimeConfig from shared.services.telemetry.api_metrics import ApiRequestTelemetryMetrics from shared.services.telemetry.aggregates import ( TelemetryAggregateSettings, @@ -26,7 +26,6 @@ stop_self_hosted_aggregate_telemetry, ) from shared.services.telemetry.events import ( - SCHEMA_VERSION, build_base_event_properties, compute_success_rate, count_to_bucket, From 896d78675f496bd8c42e02c19e7c3e1571995184 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Mon, 13 Jul 2026 22:49:39 +0800 Subject: [PATCH 05/19] Drop dashboard from OSS telemetry client allowlist. The web dashboard does not create parse jobs, so it is not an official created_by_client value. Co-authored-by: Cursor --- docs/adr/0004-anonymous-self-hosted-telemetry.md | 3 ++- packages/shared-python/shared/services/telemetry/events.py | 1 - 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/docs/adr/0004-anonymous-self-hosted-telemetry.md b/docs/adr/0004-anonymous-self-hosted-telemetry.md index d13e257a9..81a71e087 100644 --- a/docs/adr/0004-anonymous-self-hosted-telemetry.md +++ b/docs/adr/0004-anonymous-self-hosted-telemetry.md @@ -60,7 +60,8 @@ Event names use the `oss_` prefix (not `self_hosted_`). Catalog: - `document_type`: `pdf|docx|doc|xlsx|xls|pptx|ppt|csv|txt|md|html|image|other` - Extension from `job_metadata->>'source_file_name'` only; map `png|jpg|jpeg|gif|webp|tiff` → `image`; everything else → `other` -- `created_by_client`: `cli|node-sdk|dashboard|notebook|mcp|api|other` +- `created_by_client`: `cli|node-sdk|notebook|mcp|api|other` + - No `dashboard` client: the web dashboard does not create parse jobs. - From `job_metadata #>> '{document_metadata,created_by_client}'` - `source_type`: `file|url|other` diff --git a/packages/shared-python/shared/services/telemetry/events.py b/packages/shared-python/shared/services/telemetry/events.py index 533266f76..827b10b2b 100644 --- a/packages/shared-python/shared/services/telemetry/events.py +++ b/packages/shared-python/shared/services/telemetry/events.py @@ -34,7 +34,6 @@ { "cli", "node-sdk", - "dashboard", "notebook", "mcp", "api", From af90aa5e140aea3e982774f6b01b48bd999bf132 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 02:09:55 +0800 Subject: [PATCH 06/19] fix: classify dashboard jwt telemetry --- .../dashboard_jwt_authentication_service.py | 191 +++++- ...t_dashboard_jwt_authentication_contract.py | 582 ++++++++++++++++++ ...est_dashboard_token_permission_contract.py | 65 +- .../contract/test_job_creation_contract.py | 31 +- apps/api/tests/support/dashboard_jwt.py | 172 ++++++ .../core/exceptions/domain_exceptions.py | 7 +- .../core/exceptions/knowhere_exception.py | 26 +- packages/shared-python/shared/core/logging.py | 2 +- 8 files changed, 962 insertions(+), 114 deletions(-) create mode 100644 apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py create mode 100644 apps/api/tests/support/dashboard_jwt.py diff --git a/apps/api/app/services/auth/dashboard_jwt_authentication_service.py b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py index 297c68c71..52339508a 100644 --- a/apps/api/app/services/auth/dashboard_jwt_authentication_service.py +++ b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py @@ -2,23 +2,37 @@ from __future__ import annotations +import json +import re import threading from dataclasses import dataclass from datetime import timedelta -from typing import Any, Literal +from typing import Literal, cast import jwt -from jwt import PyJWKClient -from loguru import logger +from jwt import PyJWKClient, PyJWKClientConnectionError, PyJWKClientError, PyJWKSetError +from jwt.algorithms import AllowedPublicKeys from shared.core.config import settings from shared.core.exceptions.domain_exceptions import AuthException JWKS_ENDPOINT_PATH = "/api/auth/jwks" JWKS_CACHE_TTL_SECONDS = 60 * 60 +JWT_KEY_ID_MAX_LENGTH = 64 +JWT_KEY_ID_UNSAFE_PATTERN = re.compile(r"[^A-Za-z0-9._:-]") +JWT_ALGORITHMS: tuple[str, ...] = ("HS256", "RS256", "EdDSA") READ_ONLY_PERMISSION: Literal["read_only"] = "read_only" FULL_ACCESS_PERMISSION: Literal["full_access"] = "full_access" Permission = Literal["read_only", "full_access"] +JWTFailureReason = Literal[ + "jwt_missing_key_id", + "jwt_unknown_key_id", + "jwt_expired", + "jwt_invalid", + "jwks_unavailable", + "jwks_invalid", +] +VerificationKey = AllowedPublicKeys | str | bytes @dataclass(frozen=True) @@ -41,46 +55,115 @@ def decode_user_id(self, token: str) -> str: def decode_identity(self, token: str) -> DashboardJWTIdentity: """Decode and validate a JWT, returning the user ID and permission.""" try: - payload = self._decode_payload(token) + unverified_header = cast(dict[str, object], jwt.get_unverified_header(token)) + except jwt.InvalidTokenError: + raise _create_auth_exception( + user_message="Invalid token", + failure_reason="jwt_invalid", + exception_context=_build_exception_context( + algorithm=None, + key_id=None, + ), + ) from None + + algorithm = unverified_header.get("alg") + key_id_value = unverified_header.get("kid") + key_id = key_id_value if isinstance(key_id_value, str) else None + exception_context = _build_exception_context( + algorithm=algorithm, + key_id=key_id, + ) + + if key_id is None or not key_id.strip(): + raise _create_auth_exception( + failure_reason="jwt_missing_key_id", + exception_context=exception_context, + ) + + try: + key = self._get_verification_key(key_id) + if key is None: + raise _create_auth_exception( + failure_reason="jwt_unknown_key_id", + exception_context=exception_context, + ) + + payload = self._decode_payload(token, key) user_id = payload.get("id") if not isinstance(user_id, str) or not user_id: - raise AuthException(user_message="Token missing 'id' claim") + raise _create_auth_exception( + user_message="Token missing 'id' claim", + failure_reason="jwt_invalid", + exception_context=exception_context, + ) permission = _normalize_permission(payload.get("permission")) return DashboardJWTIdentity(user_id=user_id, permission=permission) except jwt.ExpiredSignatureError: - raise AuthException(user_message="Token has expired") - except jwt.InvalidTokenError as exc: - logger.warning(f"Invalid JWT token: {exc}") - raise AuthException(user_message="Invalid token") - - def _decode_payload(self, token: str) -> dict[str, Any]: - key = self._get_verification_key(token) - payload: dict[str, Any] = jwt.decode( - token, - key, - algorithms=["HS256", "RS256", "EdDSA"], - leeway=timedelta(seconds=30), - options={"verify_aud": False}, + raise _create_auth_exception( + user_message="Token has expired", + failure_reason="jwt_expired", + exception_context=exception_context, + ) from None + except PyJWKClientConnectionError: + raise _create_auth_exception( + failure_reason="jwks_unavailable", + error_category="system", + exception_context=exception_context, + ) from None + except (json.JSONDecodeError, UnicodeDecodeError, PyJWKSetError, jwt.PyJWKError): + raise _create_auth_exception( + failure_reason="jwks_invalid", + error_category="system", + exception_context=exception_context, + ) from None + except PyJWKClientError: + raise _create_auth_exception( + failure_reason="jwks_invalid", + error_category="system", + exception_context=exception_context, + ) from None + except jwt.InvalidTokenError: + raise _create_auth_exception( + user_message="Invalid token", + failure_reason="jwt_invalid", + exception_context=exception_context, + ) from None + + def _decode_payload( + self, + token: str, + key: VerificationKey, + ) -> dict[str, object]: + payload = cast( + dict[str, object], + jwt.decode( + token, + key, + algorithms=list(JWT_ALGORITHMS), + leeway=timedelta(seconds=30), + options={"verify_aud": False}, + ), ) return payload - def _get_verification_key(self, token: str) -> Any: + def _get_verification_key(self, key_id: str) -> VerificationKey | None: """Resolve the JWT verification key from the Dashboard JWKS endpoint.""" - try: - jwks_client = self._get_jwks_client() - signing_key = jwks_client.get_signing_key_from_jwt(token) - return signing_key.key - except jwt.PyJWKClientError as exc: - logger.error(f"Failed to fetch JWKS: {exc}") - raise AuthException( - internal_message=( - f"Failed to fetch verification key from JWKS endpoint: {exc}" - ) - ) - except jwt.PyJWKSetError as exc: - logger.error(f"Invalid JWKS format: {exc}") - raise AuthException(internal_message=f"Invalid JWKS format: {exc}") + jwks_client = self._get_jwks_client() + signing_keys = jwks_client.get_signing_keys() + signing_key = jwks_client.match_kid(signing_keys, key_id) + if signing_key is not None: + return cast(VerificationKey, signing_key.key) + + refreshed_signing_keys = jwks_client.get_signing_keys(refresh=True) + refreshed_signing_key = jwks_client.match_kid( + refreshed_signing_keys, + key_id, + ) + if refreshed_signing_key is None: + return None + + return cast(VerificationKey, refreshed_signing_key.key) def _get_jwks_client(self) -> PyJWKClient: """Return a cached JWKS client for Dashboard token verification.""" @@ -96,11 +179,51 @@ def _get_jwks_client(self) -> PyJWKClient: lifespan=JWKS_CACHE_TTL_SECONDS, timeout=30, ) - logger.info(f"Initialized JWKS client with endpoint: {jwks_url}") return self._jwks_client +def _build_exception_context( + *, + algorithm: object, + key_id: str | None, +) -> dict[str, object]: + is_key_id_present = key_id is not None and bool(key_id.strip()) + context: dict[str, object] = { + "auth_component": "dashboard_jwt", + "jwt_kid_present": is_key_id_present, + } + if isinstance(algorithm, str) and algorithm in JWT_ALGORITHMS: + context["jwt_algorithm"] = algorithm + if is_key_id_present and key_id is not None: + context["jwt_kid"] = _sanitize_key_id(key_id) + return context + + +def _sanitize_key_id(key_id: str) -> str: + sanitized_key_id = JWT_KEY_ID_UNSAFE_PATTERN.sub("_", key_id) + return sanitized_key_id[:JWT_KEY_ID_MAX_LENGTH] + + +def _create_auth_exception( + *, + failure_reason: JWTFailureReason, + user_message: str = "Authentication required", + error_category: Literal["client", "system"] | None = None, + exception_context: dict[str, object] | None = None, +) -> AuthException: + context: dict[str, object] = { + **(exception_context or {}), + "failure_reason": failure_reason, + } + return AuthException( + user_message=user_message, + internal_message=f"Dashboard JWT authentication failed: {failure_reason}", + error_category=error_category, + exception_context=context, + ) + + def _normalize_permission(value: object) -> Permission: if value == READ_ONLY_PERMISSION: return READ_ONLY_PERMISSION diff --git a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py new file mode 100644 index 000000000..869ddbd9a --- /dev/null +++ b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py @@ -0,0 +1,582 @@ +from __future__ import annotations + +import json +from collections.abc import Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from typing import TYPE_CHECKING, Protocol, cast + +import jwt +import pytest +from fastapi import FastAPI, Header +from httpx import ASGITransport, AsyncClient, Response +from loguru import logger +from pytest import MonkeyPatch + +from app.core.exception_handlers import setup_exception_handlers +from app.services.auth.dashboard_jwt_authentication_service import ( + DashboardJWTAuthenticationService, +) +from shared.core.config import settings +from shared.core.exceptions.domain_exceptions import AuthException +from shared.core.logging import _downgrade_expected_logfire_exception +from tests.support.dashboard_jwt import ( + create_dashboard_rsa_jwk as _create_rsa_jwk, + create_dashboard_rsa_private_key as _create_rsa_private_key, + create_dashboard_rsa_token as _create_rsa_token, + serve_dashboard_jwks as _serve_jwks, +) + +if TYPE_CHECKING: + from logfire.types import ExceptionCallbackHelper + + +class _LoguruMessage(Protocol): + @property + def record(self) -> Mapping[str, object]: ... + + +@dataclass(frozen=True) +class _CapturedAuthLog: + level: str + event: str + message: str + extra: Mapping[str, object] + + +@dataclass +class _FakeLogfireExceptionHelper: + exception: BaseException + level: str = "error" + is_recording_exception: bool = True + + def no_record_exception(self) -> None: + self.is_recording_exception = False + + +class _AuthLogCapture: + def __init__(self) -> None: + self.records: list[_CapturedAuthLog] = [] + + def capture(self, message: _LoguruMessage) -> None: + record = message.record + extra = cast(Mapping[str, object], record["extra"]) + if extra.get("auth_component") != "dashboard_jwt": + return + + level = record["level"] + self.records.append( + _CapturedAuthLog( + level=str(getattr(level, "name", level)), + event=str(extra.get("event", "")), + message=str(record["message"]), + extra=dict(extra), + ) + ) + + +@contextmanager +def _capture_auth_logs() -> Iterator[_AuthLogCapture]: + log_capture = _AuthLogCapture() + log_sink_id = logger.add(log_capture.capture) + try: + yield log_capture + finally: + logger.remove(log_sink_id) + + +def _create_authentication_app() -> FastAPI: + authentication_service = DashboardJWTAuthenticationService() + app = FastAPI() + + @app.get("/protected") + async def read_protected_resource( + authorization: str = Header(), + ) -> dict[str, str]: + _, _, token = authorization.partition(" ") + identity = authentication_service.decode_identity(token) + return { + "user_id": identity.user_id, + "permission": identity.permission, + } + + setup_exception_handlers(app) + return app + + +async def _request_with_token(token: str) -> Response: + app = _create_authentication_app() + transport = ASGITransport(app=app) + async with AsyncClient(transport=transport, base_url="http://test") as client: + return await client.get( + "/protected", + headers={"Authorization": f"Bearer {token}"}, + ) + + +def _assert_unauthenticated_response( + response: Response, + *, + expected_message: str, + token: str, +) -> None: + assert response.status_code == 401 + response_json = cast(dict[str, object], response.json()) + error = cast(dict[str, object], response_json["error"]) + assert error["code"] == "UNAUTHENTICATED" + assert error["message"] == expected_message + + serialized_response = json.dumps(response_json, default=str) + assert "failure_reason" not in serialized_response + assert "auth_component" not in serialized_response + assert "dashboard_jwt" not in serialized_response + assert "jwt_algorithm" not in serialized_response + assert "jwt_kid" not in serialized_response + assert token not in serialized_response + assert f"Bearer {token}" not in serialized_response + token_segments = token.split(".") + if len(token_segments) > 1: + assert token_segments[1] not in serialized_response + + +def _assert_log_excludes_token( + auth_log: _CapturedAuthLog, + *, + token: str, +) -> None: + serialized_log = json.dumps(auth_log.extra, default=str) + assert token not in serialized_log + assert f"Bearer {token}" not in serialized_log + token_segments = token.split(".") + if len(token_segments) > 1: + assert token_segments[1] not in serialized_log + assert "contract-dashboard-user" not in serialized_log + + +def _create_token_without_key_id() -> str: + return jwt.encode( + { + "id": "contract-dashboard-user", + "exp": datetime.now(timezone.utc) + timedelta(minutes=5), + }, + "contract-secret-with-at-least-32-bytes", + algorithm="HS256", + ) + + +@pytest.mark.asyncio +async def test_missing_key_id_is_a_client_warning_without_fetching_jwks( + monkeypatch: MonkeyPatch, +) -> None: + token = _create_token_without_key_id() + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Authentication required", + token=token, + ) + assert jwks_server.state.request_count == 0 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + assert auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + assert auth_log.extra["failure_reason"] == "jwt_missing_key_id" + assert auth_log.extra["jwt_algorithm"] == "HS256" + assert auth_log.extra["jwt_kid_present"] is False + assert "jwt_kid" not in auth_log.extra + + _assert_log_excludes_token(auth_log, token=token) + + +@pytest.mark.asyncio +async def test_malformed_jwt_is_a_client_warning_without_fetching_jwks( + monkeypatch: MonkeyPatch, +) -> None: + token = "not-a-jwt" + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Invalid token", + token=token, + ) + assert jwks_server.state.request_count == 0 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + assert auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + assert auth_log.extra["failure_reason"] == "jwt_invalid" + assert auth_log.extra["jwt_kid_present"] is False + assert "jwt_algorithm" not in auth_log.extra + assert "jwt_kid" not in auth_log.extra + _assert_log_excludes_token(auth_log, token=token) + + +@pytest.mark.asyncio +async def test_unknown_key_id_refreshes_once_and_remains_a_client_warning( + monkeypatch: MonkeyPatch, +) -> None: + signing_key = _create_rsa_private_key() + attacker_key_id = f"unknown\nkey:{'x' * 80}" + token = _create_rsa_token( + signing_key, + key_id=attacker_key_id, + ) + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_json_response( + {"keys": [_create_rsa_jwk(signing_key, key_id="known-key")]} + ) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Authentication required", + token=token, + ) + assert jwks_server.state.request_count == 2 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + assert auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + assert auth_log.extra["failure_reason"] == "jwt_unknown_key_id" + assert auth_log.extra["jwt_algorithm"] == "RS256" + assert auth_log.extra["jwt_kid_present"] is True + assert auth_log.extra["jwt_kid"] == "unknown_key:" + ("x" * 52) + + _assert_log_excludes_token(auth_log, token=token) + serialized_log = json.dumps(auth_log.extra, default=str) + assert attacker_key_id not in serialized_log + + +@pytest.mark.asyncio +async def test_unavailable_jwks_is_a_system_error_with_an_unchanged_response( + monkeypatch: MonkeyPatch, +) -> None: + signing_key = _create_rsa_private_key() + token = _create_rsa_token(signing_key, key_id="unavailable-key") + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_raw_response(b"unavailable", status_code=503) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Authentication required", + token=token, + ) + assert jwks_server.state.request_count == 1 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + assert auth_log.level == "ERROR" + assert auth_log.event == "exception.system" + assert auth_log.extra["error_category"] == "system" + assert auth_log.extra["failure_reason"] == "jwks_unavailable" + assert auth_log.extra["jwt_kid"] == "unavailable-key" + + _assert_log_excludes_token(auth_log, token=token) + + +@pytest.mark.parametrize( + "jwks_body", + [ + b'{"keys": []}', + b"not-json", + ], + ids=["empty-key-set", "malformed-json"], +) +@pytest.mark.asyncio +async def test_invalid_jwks_is_a_system_error_with_an_unchanged_response( + monkeypatch: MonkeyPatch, + jwks_body: bytes, +) -> None: + signing_key = _create_rsa_private_key() + token = _create_rsa_token(signing_key, key_id="invalid-jwks-key") + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_raw_response(jwks_body) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Authentication required", + token=token, + ) + assert jwks_server.state.request_count == 1 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + assert auth_log.level == "ERROR" + assert auth_log.event == "exception.system" + assert auth_log.extra["error_category"] == "system" + assert auth_log.extra["failure_reason"] == "jwks_invalid" + assert auth_log.extra["jwt_kid"] == "invalid-jwks-key" + + _assert_log_excludes_token(auth_log, token=token) + + +@pytest.mark.asyncio +async def test_valid_keyed_jwt_returns_identity_without_auth_rejection_log( + monkeypatch: MonkeyPatch, +) -> None: + signing_key = _create_rsa_private_key() + key_id = "valid-key" + token = _create_rsa_token( + signing_key, + key_id=key_id, + payload_overrides={"permission": "read_only"}, + ) + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_json_response( + {"keys": [_create_rsa_jwk(signing_key, key_id=key_id)]} + ) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + response = await _request_with_token(token) + + assert response.status_code == 200 + assert response.json() == { + "user_id": "contract-dashboard-user", + "permission": "read_only", + } + assert jwks_server.state.request_count == 1 + assert log_capture.records == [] + + +@pytest.mark.asyncio +async def test_expired_jwt_retains_response_and_logs_client_warning( + monkeypatch: MonkeyPatch, +) -> None: + signing_key = _create_rsa_private_key() + key_id = "expired-key" + token = _create_rsa_token( + signing_key, + key_id=key_id, + payload_overrides={"exp": datetime.now(timezone.utc) - timedelta(minutes=5)}, + ) + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_json_response( + {"keys": [_create_rsa_jwk(signing_key, key_id=key_id)]} + ) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Token has expired", + token=token, + ) + assert jwks_server.state.request_count == 1 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + assert auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + assert auth_log.extra["error_category"] == "client" + assert auth_log.extra["failure_reason"] == "jwt_expired" + assert auth_log.extra["jwt_algorithm"] == "RS256" + assert auth_log.extra["jwt_kid"] == key_id + _assert_log_excludes_token(auth_log, token=token) + + +@pytest.mark.asyncio +async def test_invalid_signature_retains_response_and_logs_client_warning( + monkeypatch: MonkeyPatch, +) -> None: + signing_key = _create_rsa_private_key() + jwks_key = _create_rsa_private_key() + key_id = "invalid-signature-key" + token = _create_rsa_token(signing_key, key_id=key_id) + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_json_response( + {"keys": [_create_rsa_jwk(jwks_key, key_id=key_id)]} + ) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Invalid token", + token=token, + ) + assert jwks_server.state.request_count == 1 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + assert auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + assert auth_log.extra["error_category"] == "client" + assert auth_log.extra["failure_reason"] == "jwt_invalid" + assert auth_log.extra["jwt_algorithm"] == "RS256" + assert auth_log.extra["jwt_kid"] == key_id + _assert_log_excludes_token(auth_log, token=token) + + +@pytest.mark.asyncio +async def test_missing_user_claim_retains_response_and_logs_client_warning( + monkeypatch: MonkeyPatch, +) -> None: + signing_key = _create_rsa_private_key() + key_id = "missing-user-key" + token = _create_rsa_token( + signing_key, + key_id=key_id, + payload_overrides={"id": None}, + ) + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_json_response( + {"keys": [_create_rsa_jwk(signing_key, key_id=key_id)]} + ) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Token missing 'id' claim", + token=token, + ) + assert jwks_server.state.request_count == 1 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + assert auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + assert auth_log.extra["error_category"] == "client" + assert auth_log.extra["failure_reason"] == "jwt_invalid" + assert auth_log.extra["jwt_algorithm"] == "RS256" + assert auth_log.extra["jwt_kid"] == key_id + _assert_log_excludes_token(auth_log, token=token) + + +@pytest.mark.asyncio +async def test_non_allowlisted_algorithm_is_not_recorded_as_safe_metadata( + monkeypatch: MonkeyPatch, +) -> None: + signing_key = _create_rsa_private_key() + key_id = "unsupported-algorithm-key" + token = _create_rsa_token( + signing_key, + key_id=key_id, + algorithm="PS256", + ) + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_json_response( + {"keys": [_create_rsa_jwk(signing_key, key_id=key_id)]} + ) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Invalid token", + token=token, + ) + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + assert auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + assert auth_log.extra["failure_reason"] == "jwt_invalid" + assert auth_log.extra["jwt_kid_present"] is True + assert auth_log.extra["jwt_kid"] == key_id + assert "jwt_algorithm" not in auth_log.extra + _assert_log_excludes_token(auth_log, token=token) + + +def test_logfire_exception_callback_keeps_system_category_auth_errors_recorded() -> None: + helper = _FakeLogfireExceptionHelper( + exception=AuthException( + error_category="system", + exception_context={ + "auth_component": "dashboard_jwt", + "failure_reason": "jwks_unavailable", + }, + ) + ) + + _downgrade_expected_logfire_exception( + cast("ExceptionCallbackHelper", helper), + ) + + assert helper.level == "error" + assert helper.is_recording_exception is True + + +def test_logfire_exception_callback_downgrades_client_category_auth_errors() -> None: + helper = _FakeLogfireExceptionHelper(exception=AuthException()) + + _downgrade_expected_logfire_exception( + cast("ExceptionCallbackHelper", helper), + ) + + assert helper.level == "warning" + assert helper.is_recording_exception is False diff --git a/apps/api/tests/contract/test_dashboard_token_permission_contract.py b/apps/api/tests/contract/test_dashboard_token_permission_contract.py index 4bd74edbe..baf1216e6 100644 --- a/apps/api/tests/contract/test_dashboard_token_permission_contract.py +++ b/apps/api/tests/contract/test_dashboard_token_permission_contract.py @@ -1,47 +1,15 @@ from collections.abc import Callable from contextlib import AbstractAsyncContextManager -from datetime import datetime, timedelta, timezone -from typing import Literal, cast +from typing import cast from uuid import uuid4 -import jwt import pytest from httpx import AsyncClient from pytest import MonkeyPatch from shared.testing.contract_runtime import seed_contract_developer from tests.support.contract_database import ContractDatabase - -Permission = Literal["read_only", "full_access"] - - -def _use_dashboard_token( - api_client: AsyncClient, - monkeypatch: MonkeyPatch, - *, - user_id: str, - permission: Permission | None, -) -> None: - jwt_secret = f"contract-jwt-secret-{uuid4().hex[:12]}" - payload: dict[str, object] = { - "id": user_id, - "exp": datetime.now(timezone.utc) + timedelta(minutes=5), - } - if permission is not None: - payload["permission"] = permission - - token = jwt.encode(payload, jwt_secret, algorithm="HS256") - - from app.services.auth.dashboard_jwt_authentication_service import ( - get_dashboard_jwt_authentication_service, - ) - - monkeypatch.setattr( - get_dashboard_jwt_authentication_service(), - "_get_verification_key", - lambda _token: jwt_secret, - ) - api_client.headers.update({"Authorization": f"Bearer {token}"}) +from tests.support.dashboard_jwt import use_dashboard_jwks_token async def _seed_dashboard_user() -> str: @@ -105,19 +73,21 @@ async def test_read_only_dashboard_token_can_read_but_cannot_parse_or_archive( user_id=user_id, namespace="contract-permission", ) - _use_dashboard_token( + + with use_dashboard_jwks_token( api_client, monkeypatch, user_id=user_id, permission="read_only", - ) - - list_jobs_response = await api_client.get("/api/v1/jobs") - get_document_response = await api_client.get(f"/api/v1/documents/{document_id}") - create_job_response = await api_client.post("/api/v1/jobs", json=payload) - archive_document_response = await api_client.post( - f"/api/v1/documents/{document_id}/archive" - ) + ): + list_jobs_response = await api_client.get("/api/v1/jobs") + get_document_response = await api_client.get( + f"/api/v1/documents/{document_id}" + ) + create_job_response = await api_client.post("/api/v1/jobs", json=payload) + archive_document_response = await api_client.post( + f"/api/v1/documents/{document_id}/archive" + ) assert list_jobs_response.status_code == 200 assert get_document_response.status_code == 200 @@ -146,14 +116,13 @@ async def test_dashboard_token_without_permission_claim_keeps_full_access( async with api_client_factory() as api_client: user_id = await _seed_dashboard_user() - _use_dashboard_token( + + with use_dashboard_jwks_token( api_client, monkeypatch, user_id=user_id, - permission=None, - ) - - response = await api_client.post("/api/v1/jobs", json=payload) + ): + response = await api_client.post("/api/v1/jobs", json=payload) assert response.status_code == 200 response_json = cast(dict[str, object], response.json()) diff --git a/apps/api/tests/contract/test_job_creation_contract.py b/apps/api/tests/contract/test_job_creation_contract.py index 5d419de87..a375f5151 100644 --- a/apps/api/tests/contract/test_job_creation_contract.py +++ b/apps/api/tests/contract/test_job_creation_contract.py @@ -1,12 +1,11 @@ from collections.abc import Callable from contextlib import AbstractAsyncContextManager -from datetime import datetime, timedelta, timezone +from datetime import datetime, timezone import json import socket from typing import cast from uuid import uuid4 -import jwt import pytest from httpx import AsyncClient from pytest import MonkeyPatch @@ -15,6 +14,7 @@ from shared.testing.contract_runtime import get_contract_database_url from tests.support.contract_database import ContractDatabase +from tests.support.dashboard_jwt import use_dashboard_jwks_token async def _create_contract_engine() -> AsyncEngine: @@ -695,15 +695,6 @@ async def test_should_reject_authenticated_user_id_missing_from_user_table( monkeypatch: MonkeyPatch, ) -> None: user_id = f"contract-missing-user-{uuid4().hex[:12]}" - jwt_secret = f"contract-jwt-secret-{uuid4().hex[:12]}" - token = jwt.encode( - { - "id": user_id, - "exp": datetime.now(timezone.utc) + timedelta(minutes=5), - }, - jwt_secret, - algorithm="HS256", - ) payload: dict[str, str] = { "namespace": "contract-jobs", "source_type": "file", @@ -712,18 +703,12 @@ async def test_should_reject_authenticated_user_id_missing_from_user_table( } async with api_client_factory() as api_client: - from app.services.auth.dashboard_jwt_authentication_service import ( - get_dashboard_jwt_authentication_service, - ) - - monkeypatch.setattr( - get_dashboard_jwt_authentication_service(), - "_get_verification_key", - lambda _token: jwt_secret, - ) - - api_client.headers.update({"Authorization": f"Bearer {token}"}) - response = await api_client.post("/api/v1/jobs", json=payload) + with use_dashboard_jwks_token( + api_client, + monkeypatch, + user_id=user_id, + ): + response = await api_client.post("/api/v1/jobs", json=payload) assert response.status_code == 401 assert response.headers["x-request-id"] diff --git a/apps/api/tests/support/dashboard_jwt.py b/apps/api/tests/support/dashboard_jwt.py new file mode 100644 index 000000000..70d1b6738 --- /dev/null +++ b/apps/api/tests/support/dashboard_jwt.py @@ -0,0 +1,172 @@ +from __future__ import annotations + +import json +import threading +from collections.abc import Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass, field +from datetime import datetime, timedelta, timezone +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import Literal, cast + +import jwt +from cryptography.hazmat.primitives.asymmetric.rsa import ( + RSAPrivateKey, + generate_private_key, +) +from httpx import AsyncClient +from pytest import MonkeyPatch + +DashboardPermission = Literal["read_only", "full_access"] + + +@dataclass +class JWKSResponseState: + body: bytes = b'{"keys": []}' + status_code: int = 200 + request_count: int = 0 + lock: threading.Lock = field(default_factory=threading.Lock) + + def record_request(self) -> tuple[bytes, int]: + with self.lock: + self.request_count += 1 + return self.body, self.status_code + + def set_json_response( + self, + response: Mapping[str, object], + *, + status_code: int = 200, + ) -> None: + with self.lock: + self.body = json.dumps(response).encode("utf-8") + self.status_code = status_code + + def set_raw_response( + self, + body: bytes, + *, + status_code: int = 200, + ) -> None: + with self.lock: + self.body = body + self.status_code = status_code + + +@dataclass(frozen=True) +class LocalJWKSServer: + endpoint: str + state: JWKSResponseState + + +def create_dashboard_rsa_private_key() -> RSAPrivateKey: + return generate_private_key(public_exponent=65537, key_size=2048) + + +def create_dashboard_rsa_jwk( + private_key: RSAPrivateKey, + *, + key_id: str, +) -> dict[str, object]: + jwk = cast( + dict[str, object], + jwt.algorithms.RSAAlgorithm.to_jwk(private_key.public_key(), as_dict=True), + ) + return { + **jwk, + "kid": key_id, + "use": "sig", + "alg": "RS256", + } + + +def create_dashboard_rsa_token( + private_key: RSAPrivateKey, + *, + key_id: str, + user_id: str = "contract-dashboard-user", + permission: DashboardPermission | None = None, + expires_at: datetime | None = None, + payload_overrides: Mapping[str, object] | None = None, + algorithm: str = "RS256", +) -> str: + payload: dict[str, object] = { + "id": user_id, + "exp": expires_at or datetime.now(timezone.utc) + timedelta(minutes=5), + } + if permission is not None: + payload["permission"] = permission + payload.update(payload_overrides or {}) + + return jwt.encode( + payload, + private_key, + algorithm=algorithm, + headers={"kid": key_id}, + ) + + +def _create_jwks_handler( + state: JWKSResponseState, +) -> type[BaseHTTPRequestHandler]: + class JWKSRequestHandler(BaseHTTPRequestHandler): + def do_GET(self) -> None: + body, status_code = state.record_request() + self.send_response(status_code) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def log_message(self, format: str, *args: object) -> None: + return + + return JWKSRequestHandler + + +@contextmanager +def serve_dashboard_jwks() -> Iterator[LocalJWKSServer]: + state = JWKSResponseState() + server = ThreadingHTTPServer(("127.0.0.1", 0), _create_jwks_handler(state)) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + host, port = server.server_address + + try: + yield LocalJWKSServer(endpoint=f"http://{host}:{port}", state=state) + finally: + server.shutdown() + server.server_close() + thread.join(timeout=5) + + +@contextmanager +def use_dashboard_jwks_token( + api_client: AsyncClient, + monkeypatch: MonkeyPatch, + *, + user_id: str, + permission: DashboardPermission | None = None, +) -> Iterator[str]: + key_id = f"contract-dashboard-key-{user_id}" + signing_key = create_dashboard_rsa_private_key() + token = create_dashboard_rsa_token( + signing_key, + key_id=key_id, + user_id=user_id, + permission=permission, + ) + + with serve_dashboard_jwks() as jwks_server: + from shared.core.config import settings + + jwks_server.state.set_json_response( + {"keys": [create_dashboard_rsa_jwk(signing_key, key_id=key_id)]} + ) + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + jwks_server.endpoint, + ) + api_client.headers.update({"Authorization": f"Bearer {token}"}) + yield token diff --git a/packages/shared-python/shared/core/exceptions/domain_exceptions.py b/packages/shared-python/shared/core/exceptions/domain_exceptions.py index 57ddd79a1..6b8b518be 100644 --- a/packages/shared-python/shared/core/exceptions/domain_exceptions.py +++ b/packages/shared-python/shared/core/exceptions/domain_exceptions.py @@ -42,9 +42,10 @@ # User sees: "An internal system error occurred. Please contact support." """ +from collections.abc import Mapping from typing import Any, Dict, List, Optional, TypedDict -from shared.core.exceptions.knowhere_exception import KnowhereException +from shared.core.exceptions.knowhere_exception import ErrorCategory, KnowhereException from shared.core.response.ErrorCode import ErrorCode, SubCode # ============================================================================ @@ -112,12 +113,16 @@ def __init__( self, user_message: str = "Authentication required", internal_message: Optional[str] = None, + error_category: ErrorCategory | None = None, + exception_context: Mapping[str, object] | None = None, ): super().__init__( code=ErrorCode.UNAUTHENTICATED, internal_message=internal_message or user_message, user_message=user_message, details={}, # Empty for security + error_category=error_category, + exception_context=exception_context, ) diff --git a/packages/shared-python/shared/core/exceptions/knowhere_exception.py b/packages/shared-python/shared/core/exceptions/knowhere_exception.py index 9158be63f..22c91fbc3 100644 --- a/packages/shared-python/shared/core/exceptions/knowhere_exception.py +++ b/packages/shared-python/shared/core/exceptions/knowhere_exception.py @@ -59,13 +59,15 @@ raise KnowhereException(code=ErrorCode.INVALID_ARGUMENT, ...) """ -from typing import Any, Dict, Optional +from collections.abc import Mapping +from typing import Any, Dict, Literal, Optional from shared.core.response.ErrorCode import ErrorCode, ErrorCodeMapper # Default messages for auto-sanitization DEFAULT_5XX_USER_MESSAGE = "An internal system error occurred. Please contact support." DEFAULT_4XX_USER_MESSAGE = "Invalid request. Please check your input." +ErrorCategory = Literal["client", "system"] class KnowhereException(Exception): @@ -121,6 +123,8 @@ def __init__( details: Optional[Dict[str, Any]] = None, http_status_code: Optional[int] = None, original_exception: Optional[Exception] = None, + error_category: ErrorCategory | None = None, + exception_context: Mapping[str, object] | None = None, ): """ Initialize a KnowhereException. @@ -134,6 +138,10 @@ def __init__( details: Optional structured data to include in response (must be safe). http_status_code: Override HTTP status (auto-derived from code if None). original_exception: The underlying exception being wrapped (for logging). + error_category: Optional telemetry category override. Defaults to the + category derived from the HTTP status. + exception_context: Internal-only structured telemetry fields. These are + included in logs and never returned to clients. """ super().__init__(internal_message) self.code = code @@ -143,6 +151,11 @@ def __init__( http_status_code or ErrorCodeMapper.get_http_status_from_error_code(code) ) self.original_exception = original_exception + default_error_category: ErrorCategory = ( + "system" if self.http_status_code >= 500 else "client" + ) + self.error_category: ErrorCategory = error_category or default_error_category + self.exception_context: Dict[str, object] = dict(exception_context or {}) # ======================================================================= # SECURITY: Auto-sanitize user_message based on HTTP status @@ -215,13 +228,11 @@ def to_log(self) -> Dict[str, Any]: - details: Additional structured data - original_exception: Wrapped exception info """ - # Determine error category based on HTTP status - error_category = "system" if self.http_status_code >= 500 else "client" - log_data: Dict[str, Any] = { + **self.exception_context, "error_code": self.code.value, "http_status": self.http_status_code, - "error_category": error_category, + "error_category": self.error_category, "exception_class": self.__class__.__name__, "internal_message": self.internal_message, "user_message": self.user_message, @@ -275,8 +286,9 @@ def logging(self, **extra_context): } # Log at appropriate level with appropriate event - if self.http_status_code >= 500: - # 5xx: ERROR level with stacktrace + if self.error_category == "system": + # System-category errors use ERROR level with a stacktrace even when + # their public HTTP status intentionally remains a 4xx response. logger.bind(event=LogEvent.EXCEPTION_SYSTEM.value, **log_data).opt( exception=self ).error(self.internal_message) diff --git a/packages/shared-python/shared/core/logging.py b/packages/shared-python/shared/core/logging.py index bb246f18e..60b815666 100644 --- a/packages/shared-python/shared/core/logging.py +++ b/packages/shared-python/shared/core/logging.py @@ -84,7 +84,7 @@ def _is_expected_client_exception(exception: BaseException) -> bool: from shared.core.exceptions.knowhere_exception import KnowhereException if isinstance(exception, KnowhereException): - return 400 <= exception.http_status_code < 500 + return exception.error_category == "client" try: from fastapi import HTTPException as FastAPIHTTPException From abd83fbde7777d6f59b24a194a734f5f222368f7 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 02:26:37 +0800 Subject: [PATCH 07/19] test: address jwt telemetry code scanning comments --- ...est_dashboard_jwt_authentication_contract.py | 17 +++++------------ 1 file changed, 5 insertions(+), 12 deletions(-) diff --git a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py index 869ddbd9a..3351f5a31 100644 --- a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py +++ b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py @@ -5,7 +5,7 @@ from contextlib import contextmanager from dataclasses import dataclass from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Protocol, cast +from typing import Protocol, cast import jwt import pytest @@ -28,13 +28,10 @@ serve_dashboard_jwks as _serve_jwks, ) -if TYPE_CHECKING: - from logfire.types import ExceptionCallbackHelper - - class _LoguruMessage(Protocol): @property - def record(self) -> Mapping[str, object]: ... + def record(self) -> Mapping[str, object]: + raise NotImplementedError @dataclass(frozen=True) @@ -563,9 +560,7 @@ def test_logfire_exception_callback_keeps_system_category_auth_errors_recorded() ) ) - _downgrade_expected_logfire_exception( - cast("ExceptionCallbackHelper", helper), - ) + _downgrade_expected_logfire_exception(helper) # pyright: ignore[reportArgumentType] assert helper.level == "error" assert helper.is_recording_exception is True @@ -574,9 +569,7 @@ def test_logfire_exception_callback_keeps_system_category_auth_errors_recorded() def test_logfire_exception_callback_downgrades_client_category_auth_errors() -> None: helper = _FakeLogfireExceptionHelper(exception=AuthException()) - _downgrade_expected_logfire_exception( - cast("ExceptionCallbackHelper", helper), - ) + _downgrade_expected_logfire_exception(helper) # pyright: ignore[reportArgumentType] assert helper.level == "warning" assert helper.is_recording_exception is False From 8e555022725f42a303674a9e13ecc100b3bfec23 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 10:11:22 +0800 Subject: [PATCH 08/19] test: isolate dashboard jwt contract imports --- ...t_dashboard_jwt_authentication_contract.py | 135 ++++++++++-------- 1 file changed, 74 insertions(+), 61 deletions(-) diff --git a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py index 3351f5a31..f9de2cf68 100644 --- a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py +++ b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py @@ -1,11 +1,13 @@ from __future__ import annotations import json +import sys from collections.abc import Iterator, Mapping from contextlib import contextmanager from dataclasses import dataclass from datetime import datetime, timedelta, timezone -from typing import Protocol, cast +from pathlib import Path +from typing import Literal, Protocol, cast import jwt import pytest @@ -14,13 +16,10 @@ from loguru import logger from pytest import MonkeyPatch -from app.core.exception_handlers import setup_exception_handlers -from app.services.auth.dashboard_jwt_authentication_service import ( - DashboardJWTAuthenticationService, +from tests.support.import_environment import ( + configure_import_environment, + ensure_import_paths, ) -from shared.core.config import settings -from shared.core.exceptions.domain_exceptions import AuthException -from shared.core.logging import _downgrade_expected_logfire_exception from tests.support.dashboard_jwt import ( create_dashboard_rsa_jwk as _create_rsa_jwk, create_dashboard_rsa_private_key as _create_rsa_private_key, @@ -28,6 +27,10 @@ serve_dashboard_jwks as _serve_jwks, ) +configure_import_environment() +ensure_import_paths() + + class _LoguruMessage(Protocol): @property def record(self) -> Mapping[str, object]: @@ -83,7 +86,57 @@ def _capture_auth_logs() -> Iterator[_AuthLogCapture]: logger.remove(log_sink_id) +def _prepare_api_app_imports() -> None: + api_root = str(Path(__file__).resolve().parents[2]) + if api_root in sys.path: + sys.path.remove(api_root) + sys.path.insert(0, api_root) + + for module_name in list(sys.modules): + if module_name == "app" or module_name.startswith("app."): + sys.modules.pop(module_name, None) + + +def _use_dashboard_endpoint( + monkeypatch: MonkeyPatch, + endpoint: str, +) -> None: + from shared.core.config import settings + + monkeypatch.setattr( + settings, + "INTERNAL_DASHBOARD_ENDPOINT", + endpoint, + ) + + +def _create_auth_exception( + *, + error_category: Literal["client", "system"] | None = None, + exception_context: Mapping[str, object] | None = None, +) -> BaseException: + from shared.core.exceptions.domain_exceptions import AuthException + + return AuthException( + error_category=error_category, + exception_context=exception_context, + ) + + +def _downgrade_logfire_exception(helper: _FakeLogfireExceptionHelper) -> None: + from shared.core.logging import _downgrade_expected_logfire_exception + + _downgrade_expected_logfire_exception(helper) # pyright: ignore[reportArgumentType] + + def _create_authentication_app() -> FastAPI: + _prepare_api_app_imports() + + from app.core.exception_handlers import setup_exception_handlers + from app.services.auth.dashboard_jwt_authentication_service import ( + DashboardJWTAuthenticationService, + ) + authentication_service = DashboardJWTAuthenticationService() app = FastAPI() @@ -170,11 +223,7 @@ async def test_missing_key_id_is_a_client_warning_without_fetching_jwks( with _capture_auth_logs() as log_capture: with _serve_jwks() as jwks_server: - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -204,11 +253,7 @@ async def test_malformed_jwt_is_a_client_warning_without_fetching_jwks( with _capture_auth_logs() as log_capture: with _serve_jwks() as jwks_server: - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -245,11 +290,7 @@ async def test_unknown_key_id_refreshes_once_and_remains_a_client_warning( jwks_server.state.set_json_response( {"keys": [_create_rsa_jwk(signing_key, key_id="known-key")]} ) - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -283,11 +324,7 @@ async def test_unavailable_jwks_is_a_system_error_with_an_unchanged_response( with _capture_auth_logs() as log_capture: with _serve_jwks() as jwks_server: jwks_server.state.set_raw_response(b"unavailable", status_code=503) - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -327,11 +364,7 @@ async def test_invalid_jwks_is_a_system_error_with_an_unchanged_response( with _capture_auth_logs() as log_capture: with _serve_jwks() as jwks_server: jwks_server.state.set_raw_response(jwks_body) - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -369,11 +402,7 @@ async def test_valid_keyed_jwt_returns_identity_without_auth_rejection_log( jwks_server.state.set_json_response( {"keys": [_create_rsa_jwk(signing_key, key_id=key_id)]} ) - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) assert response.status_code == 200 @@ -402,11 +431,7 @@ async def test_expired_jwt_retains_response_and_logs_client_warning( jwks_server.state.set_json_response( {"keys": [_create_rsa_jwk(signing_key, key_id=key_id)]} ) - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -441,11 +466,7 @@ async def test_invalid_signature_retains_response_and_logs_client_warning( jwks_server.state.set_json_response( {"keys": [_create_rsa_jwk(jwks_key, key_id=key_id)]} ) - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -483,11 +504,7 @@ async def test_missing_user_claim_retains_response_and_logs_client_warning( jwks_server.state.set_json_response( {"keys": [_create_rsa_jwk(signing_key, key_id=key_id)]} ) - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -525,11 +542,7 @@ async def test_non_allowlisted_algorithm_is_not_recorded_as_safe_metadata( jwks_server.state.set_json_response( {"keys": [_create_rsa_jwk(signing_key, key_id=key_id)]} ) - monkeypatch.setattr( - settings, - "INTERNAL_DASHBOARD_ENDPOINT", - jwks_server.endpoint, - ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) _assert_unauthenticated_response( @@ -551,7 +564,7 @@ async def test_non_allowlisted_algorithm_is_not_recorded_as_safe_metadata( def test_logfire_exception_callback_keeps_system_category_auth_errors_recorded() -> None: helper = _FakeLogfireExceptionHelper( - exception=AuthException( + exception=_create_auth_exception( error_category="system", exception_context={ "auth_component": "dashboard_jwt", @@ -560,16 +573,16 @@ def test_logfire_exception_callback_keeps_system_category_auth_errors_recorded() ) ) - _downgrade_expected_logfire_exception(helper) # pyright: ignore[reportArgumentType] + _downgrade_logfire_exception(helper) assert helper.level == "error" assert helper.is_recording_exception is True def test_logfire_exception_callback_downgrades_client_category_auth_errors() -> None: - helper = _FakeLogfireExceptionHelper(exception=AuthException()) + helper = _FakeLogfireExceptionHelper(exception=_create_auth_exception()) - _downgrade_expected_logfire_exception(helper) # pyright: ignore[reportArgumentType] + _downgrade_logfire_exception(helper) assert helper.level == "warning" assert helper.is_recording_exception is False From 9645eb5552efd95f2338b708792c7b356b56e245 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 10:21:02 +0800 Subject: [PATCH 09/19] test: type dashboard jwt import isolation helper --- .../contract/test_dashboard_jwt_authentication_contract.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py index f9de2cf68..f38320e05 100644 --- a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py +++ b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py @@ -87,12 +87,13 @@ def _capture_auth_logs() -> Iterator[_AuthLogCapture]: def _prepare_api_app_imports() -> None: - api_root = str(Path(__file__).resolve().parents[2]) + api_root: str = str(Path(__file__).resolve().parents[2]) if api_root in sys.path: sys.path.remove(api_root) sys.path.insert(0, api_root) - for module_name in list(sys.modules): + cached_module_names: list[str] = list(sys.modules) + for module_name in cached_module_names: if module_name == "app" or module_name.startswith("app."): sys.modules.pop(module_name, None) From 1cd8a011a5a8794e5e0947eebbdaf1bb5f4f35ce Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 14:46:42 +0800 Subject: [PATCH 10/19] fix: reject malformed jwt payloads before jwks lookup --- .../dashboard_jwt_authentication_service.py | 32 ++++++++++++++ ...t_dashboard_jwt_authentication_contract.py | 44 +++++++++++++++++++ 2 files changed, 76 insertions(+) diff --git a/apps/api/app/services/auth/dashboard_jwt_authentication_service.py b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py index 52339508a..518634e7f 100644 --- a/apps/api/app/services/auth/dashboard_jwt_authentication_service.py +++ b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py @@ -12,6 +12,7 @@ import jwt from jwt import PyJWKClient, PyJWKClientConnectionError, PyJWKClientError, PyJWKSetError from jwt.algorithms import AllowedPublicKeys +from jwt.types import Options from shared.core.config import settings from shared.core.exceptions.domain_exceptions import AuthException @@ -21,6 +22,16 @@ JWT_KEY_ID_MAX_LENGTH = 64 JWT_KEY_ID_UNSAFE_PATTERN = re.compile(r"[^A-Za-z0-9._:-]") JWT_ALGORITHMS: tuple[str, ...] = ("HS256", "RS256", "EdDSA") +JWT_STRUCTURE_ONLY_DECODE_OPTIONS: Options = { + "verify_signature": False, + "verify_exp": False, + "verify_nbf": False, + "verify_iat": False, + "verify_aud": False, + "verify_iss": False, + "verify_sub": False, + "verify_jti": False, +} READ_ONLY_PERMISSION: Literal["read_only"] = "read_only" FULL_ACCESS_PERMISSION: Literal["full_access"] = "full_access" Permission = Literal["read_only", "full_access"] @@ -81,6 +92,7 @@ def decode_identity(self, token: str) -> DashboardJWTIdentity: ) try: + self._reject_malformed_token_before_jwks_lookup(token, exception_context) key = self._get_verification_key(key_id) if key is None: raise _create_auth_exception( @@ -147,6 +159,26 @@ def _decode_payload( ) return payload + def _reject_malformed_token_before_jwks_lookup( + self, + token: str, + exception_context: dict[str, object], + ) -> None: + """Reject structurally invalid JWTs before touching Dashboard JWKS.""" + try: + # This decode only checks token structure; verified claims come from + # _decode_payload after the signing key is resolved. + jwt.decode( + token, + options=JWT_STRUCTURE_ONLY_DECODE_OPTIONS, + ) + except (json.JSONDecodeError, UnicodeDecodeError, jwt.InvalidTokenError): + raise _create_auth_exception( + user_message="Invalid token", + failure_reason="jwt_invalid", + exception_context=exception_context, + ) from None + def _get_verification_key(self, key_id: str) -> VerificationKey | None: """Resolve the JWT verification key from the Dashboard JWKS endpoint.""" jwks_client = self._get_jwks_client() diff --git a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py index f38320e05..4d0bc5097 100644 --- a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py +++ b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py @@ -1,5 +1,6 @@ from __future__ import annotations +import base64 import json import sys from collections.abc import Iterator, Mapping @@ -216,6 +217,18 @@ def _create_token_without_key_id() -> str: ) +def _base64url_encode(value: bytes) -> str: + return base64.urlsafe_b64encode(value).rstrip(b"=").decode("ascii") + + +def _create_token_with_malformed_json_payload(*, key_id: str) -> str: + header: dict[str, object] = {"alg": "RS256", "kid": key_id, "typ": "JWT"} + header_segment = _base64url_encode(json.dumps(header).encode("utf-8")) + payload_segment = _base64url_encode(b"not-json") + signature_segment = _base64url_encode(b"signature") + return f"{header_segment}.{payload_segment}.{signature_segment}" + + @pytest.mark.asyncio async def test_missing_key_id_is_a_client_warning_without_fetching_jwks( monkeypatch: MonkeyPatch, @@ -275,6 +288,37 @@ async def test_malformed_jwt_is_a_client_warning_without_fetching_jwks( _assert_log_excludes_token(auth_log, token=token) +@pytest.mark.asyncio +async def test_malformed_jwt_payload_is_a_client_warning_without_fetching_jwks( + monkeypatch: MonkeyPatch, +) -> None: + key_id = "malformed-payload-key" + token = _create_token_with_malformed_json_payload(key_id=key_id) + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + jwks_server.state.set_raw_response(b"unavailable", status_code=503) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Invalid token", + token=token, + ) + assert jwks_server.state.request_count == 0 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + assert auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + assert auth_log.extra["failure_reason"] == "jwt_invalid" + assert auth_log.extra["jwt_algorithm"] == "RS256" + assert auth_log.extra["jwt_kid_present"] is True + assert auth_log.extra["jwt_kid"] == key_id + _assert_log_excludes_token(auth_log, token=token) + + @pytest.mark.asyncio async def test_unknown_key_id_refreshes_once_and_remains_a_client_warning( monkeypatch: MonkeyPatch, From e7c246358246a243111c4d17c1d0d7169a02c9bd Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 16:25:25 +0800 Subject: [PATCH 11/19] fix: keep dashboard jwt telemetry local --- .../dashboard_jwt_authentication_service.py | 148 +++++++++++------- ...t_dashboard_jwt_authentication_contract.py | 143 +++++++++-------- .../core/exceptions/domain_exceptions.py | 7 +- .../core/exceptions/knowhere_exception.py | 30 ++-- packages/shared-python/shared/core/logging.py | 2 +- 5 files changed, 187 insertions(+), 143 deletions(-) diff --git a/apps/api/app/services/auth/dashboard_jwt_authentication_service.py b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py index 518634e7f..300109e3d 100644 --- a/apps/api/app/services/auth/dashboard_jwt_authentication_service.py +++ b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py @@ -7,15 +7,17 @@ import threading from dataclasses import dataclass from datetime import timedelta -from typing import Literal, cast +from typing import Literal, NoReturn, cast import jwt from jwt import PyJWKClient, PyJWKClientConnectionError, PyJWKClientError, PyJWKSetError from jwt.algorithms import AllowedPublicKeys from jwt.types import Options +from loguru import logger from shared.core.config import settings from shared.core.exceptions.domain_exceptions import AuthException +from shared.core.logging import LogEvent JWKS_ENDPOINT_PATH = "/api/auth/jwks" JWKS_CACHE_TTL_SECONDS = 60 * 60 @@ -68,79 +70,80 @@ def decode_identity(self, token: str) -> DashboardJWTIdentity: try: unverified_header = cast(dict[str, object], jwt.get_unverified_header(token)) except jwt.InvalidTokenError: - raise _create_auth_exception( - user_message="Invalid token", + _reject_client_jwt( failure_reason="jwt_invalid", - exception_context=_build_exception_context( + telemetry_context=_build_telemetry_context( algorithm=None, key_id=None, ), - ) from None + ) algorithm = unverified_header.get("alg") key_id_value = unverified_header.get("kid") key_id = key_id_value if isinstance(key_id_value, str) else None - exception_context = _build_exception_context( + telemetry_context = _build_telemetry_context( algorithm=algorithm, key_id=key_id, ) if key_id is None or not key_id.strip(): - raise _create_auth_exception( + _reject_client_jwt( failure_reason="jwt_missing_key_id", - exception_context=exception_context, + telemetry_context=telemetry_context, ) try: - self._reject_malformed_token_before_jwks_lookup(token, exception_context) + self._reject_malformed_token_before_jwks_lookup(token, telemetry_context) key = self._get_verification_key(key_id) if key is None: - raise _create_auth_exception( + _reject_client_jwt( failure_reason="jwt_unknown_key_id", - exception_context=exception_context, + telemetry_context=telemetry_context, ) payload = self._decode_payload(token, key) user_id = payload.get("id") if not isinstance(user_id, str) or not user_id: - raise _create_auth_exception( - user_message="Token missing 'id' claim", + _reject_client_jwt( failure_reason="jwt_invalid", - exception_context=exception_context, + telemetry_context=telemetry_context, ) permission = _normalize_permission(payload.get("permission")) return DashboardJWTIdentity(user_id=user_id, permission=permission) except jwt.ExpiredSignatureError: - raise _create_auth_exception( - user_message="Token has expired", + _reject_client_jwt( failure_reason="jwt_expired", - exception_context=exception_context, - ) from None - except PyJWKClientConnectionError: - raise _create_auth_exception( + telemetry_context=telemetry_context, + ) + except PyJWKClientConnectionError as error: + _reject_jwks_dependency( failure_reason="jwks_unavailable", - error_category="system", - exception_context=exception_context, - ) from None - except (json.JSONDecodeError, UnicodeDecodeError, PyJWKSetError, jwt.PyJWKError): - raise _create_auth_exception( + telemetry_context=telemetry_context, + original_exception=error, + ) + except ( + json.JSONDecodeError, + UnicodeDecodeError, + PyJWKSetError, + jwt.PyJWKError, + ) as error: + _reject_jwks_dependency( failure_reason="jwks_invalid", - error_category="system", - exception_context=exception_context, - ) from None - except PyJWKClientError: - raise _create_auth_exception( + telemetry_context=telemetry_context, + original_exception=error, + ) + except PyJWKClientError as error: + _reject_jwks_dependency( failure_reason="jwks_invalid", - error_category="system", - exception_context=exception_context, - ) from None + telemetry_context=telemetry_context, + original_exception=error, + ) except jwt.InvalidTokenError: - raise _create_auth_exception( - user_message="Invalid token", + _reject_client_jwt( failure_reason="jwt_invalid", - exception_context=exception_context, - ) from None + telemetry_context=telemetry_context, + ) def _decode_payload( self, @@ -162,7 +165,7 @@ def _decode_payload( def _reject_malformed_token_before_jwks_lookup( self, token: str, - exception_context: dict[str, object], + telemetry_context: dict[str, object], ) -> None: """Reject structurally invalid JWTs before touching Dashboard JWKS.""" try: @@ -173,11 +176,10 @@ def _reject_malformed_token_before_jwks_lookup( options=JWT_STRUCTURE_ONLY_DECODE_OPTIONS, ) except (json.JSONDecodeError, UnicodeDecodeError, jwt.InvalidTokenError): - raise _create_auth_exception( - user_message="Invalid token", + _reject_client_jwt( failure_reason="jwt_invalid", - exception_context=exception_context, - ) from None + telemetry_context=telemetry_context, + ) def _get_verification_key(self, key_id: str) -> VerificationKey | None: """Resolve the JWT verification key from the Dashboard JWKS endpoint.""" @@ -215,7 +217,7 @@ def _get_jwks_client(self) -> PyJWKClient: return self._jwks_client -def _build_exception_context( +def _build_telemetry_context( *, algorithm: object, key_id: str | None, @@ -237,23 +239,57 @@ def _sanitize_key_id(key_id: str) -> str: return sanitized_key_id[:JWT_KEY_ID_MAX_LENGTH] -def _create_auth_exception( +def _reject_client_jwt( *, failure_reason: JWTFailureReason, - user_message: str = "Authentication required", - error_category: Literal["client", "system"] | None = None, - exception_context: dict[str, object] | None = None, -) -> AuthException: - context: dict[str, object] = { - **(exception_context or {}), + telemetry_context: dict[str, object], +) -> NoReturn: + _log_dashboard_jwt_auth_failure( + failure_reason=failure_reason, + telemetry_context=telemetry_context, + is_jwks_dependency_failure=False, + ) + raise AuthException() from None + + +def _reject_jwks_dependency( + *, + failure_reason: JWTFailureReason, + telemetry_context: dict[str, object], + original_exception: Exception, +) -> NoReturn: + _log_dashboard_jwt_auth_failure( + failure_reason=failure_reason, + telemetry_context=telemetry_context, + is_jwks_dependency_failure=True, + original_exception=original_exception, + ) + raise AuthException() from None + + +def _log_dashboard_jwt_auth_failure( + *, + failure_reason: JWTFailureReason, + telemetry_context: dict[str, object], + is_jwks_dependency_failure: bool, + original_exception: Exception | None = None, +) -> None: + log_data: dict[str, object] = { + **telemetry_context, "failure_reason": failure_reason, } - return AuthException( - user_message=user_message, - internal_message=f"Dashboard JWT authentication failed: {failure_reason}", - error_category=error_category, - exception_context=context, - ) + message = f"Dashboard JWT authentication failed: {failure_reason}" + if is_jwks_dependency_failure: + logger.bind( + event=LogEvent.EXCEPTION_SYSTEM.value, + **log_data, + ).opt(exception=original_exception).error(message) + return + + logger.bind( + event=LogEvent.EXCEPTION_CLIENT.value, + **log_data, + ).warning(message) def _normalize_permission(value: object) -> Permission: diff --git a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py index 4d0bc5097..5d38d3a5a 100644 --- a/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py +++ b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py @@ -8,7 +8,7 @@ from dataclasses import dataclass from datetime import datetime, timedelta, timezone from pathlib import Path -from typing import Literal, Protocol, cast +from typing import Protocol, cast import jwt import pytest @@ -44,6 +44,8 @@ class _CapturedAuthLog: event: str message: str extra: Mapping[str, object] + exception_type: str | None + exception_message: str | None @dataclass @@ -66,13 +68,25 @@ def capture(self, message: _LoguruMessage) -> None: if extra.get("auth_component") != "dashboard_jwt": return - level = record["level"] + level: object = record["level"] + exception: object | None = record.get("exception") + exception_type: str | None = None + exception_message: str | None = None + if exception is not None: + exception_value: object | None = getattr(exception, "value", None) + exception_type_value: object | None = getattr(exception, "type", None) + if exception_type_value is not None: + exception_type = str(getattr(exception_type_value, "__name__", "")) + if exception_value is not None: + exception_message = str(exception_value) self.records.append( _CapturedAuthLog( level=str(getattr(level, "name", level)), event=str(extra.get("event", "")), message=str(record["message"]), extra=dict(extra), + exception_type=exception_type, + exception_message=exception_message, ) ) @@ -112,17 +126,10 @@ def _use_dashboard_endpoint( ) -def _create_auth_exception( - *, - error_category: Literal["client", "system"] | None = None, - exception_context: Mapping[str, object] | None = None, -) -> BaseException: +def _create_auth_exception() -> BaseException: from shared.core.exceptions.domain_exceptions import AuthException - return AuthException( - error_category=error_category, - exception_context=exception_context, - ) + return AuthException() def _downgrade_logfire_exception(helper: _FakeLogfireExceptionHelper) -> None: @@ -185,6 +192,8 @@ def _assert_unauthenticated_response( assert "dashboard_jwt" not in serialized_response assert "jwt_algorithm" not in serialized_response assert "jwt_kid" not in serialized_response + assert "payload" not in serialized_response + assert "contract-dashboard-user" not in serialized_response assert token not in serialized_response assert f"Bearer {token}" not in serialized_response token_segments = token.split(".") @@ -197,7 +206,7 @@ def _assert_log_excludes_token( *, token: str, ) -> None: - serialized_log = json.dumps(auth_log.extra, default=str) + serialized_log = _serialize_auth_log(auth_log) assert token not in serialized_log assert f"Bearer {token}" not in serialized_log token_segments = token.split(".") @@ -206,6 +215,44 @@ def _assert_log_excludes_token( assert "contract-dashboard-user" not in serialized_log +def _serialize_auth_log(auth_log: _CapturedAuthLog) -> str: + return json.dumps( + { + "message": auth_log.message, + "extra": auth_log.extra, + "exception_type": auth_log.exception_type, + "exception_message": auth_log.exception_message, + }, + default=str, + ) + + +def _assert_client_auth_log(auth_log: _CapturedAuthLog) -> None: + assert auth_log.level == "WARNING" + assert auth_log.event == "exception.client" + assert auth_log.exception_type is None + assert auth_log.exception_message is None + assert "error_category" not in auth_log.extra + + +def _assert_jwks_dependency_auth_log(auth_log: _CapturedAuthLog) -> None: + assert auth_log.level == "ERROR" + assert auth_log.event == "exception.system" + assert auth_log.exception_type is not None + assert auth_log.exception_message is not None + assert "error_category" not in auth_log.extra + + +def _assert_log_excludes_jwks_body( + auth_log: _CapturedAuthLog, + *, + jwks_body: bytes, +) -> None: + raw_body = jwks_body.decode("utf-8", errors="ignore") + if raw_body: + assert raw_body not in _serialize_auth_log(auth_log) + + def _create_token_without_key_id() -> str: return jwt.encode( { @@ -249,8 +296,7 @@ async def test_missing_key_id_is_a_client_warning_without_fetching_jwks( assert len(log_capture.records) == 1 auth_log = log_capture.records[0] - assert auth_log.level == "WARNING" - assert auth_log.event == "exception.client" + _assert_client_auth_log(auth_log) assert auth_log.extra["failure_reason"] == "jwt_missing_key_id" assert auth_log.extra["jwt_algorithm"] == "HS256" assert auth_log.extra["jwt_kid_present"] is False @@ -272,15 +318,14 @@ async def test_malformed_jwt_is_a_client_warning_without_fetching_jwks( _assert_unauthenticated_response( response, - expected_message="Invalid token", + expected_message="Authentication required", token=token, ) assert jwks_server.state.request_count == 0 assert len(log_capture.records) == 1 auth_log = log_capture.records[0] - assert auth_log.level == "WARNING" - assert auth_log.event == "exception.client" + _assert_client_auth_log(auth_log) assert auth_log.extra["failure_reason"] == "jwt_invalid" assert auth_log.extra["jwt_kid_present"] is False assert "jwt_algorithm" not in auth_log.extra @@ -303,15 +348,14 @@ async def test_malformed_jwt_payload_is_a_client_warning_without_fetching_jwks( _assert_unauthenticated_response( response, - expected_message="Invalid token", + expected_message="Authentication required", token=token, ) assert jwks_server.state.request_count == 0 assert len(log_capture.records) == 1 auth_log = log_capture.records[0] - assert auth_log.level == "WARNING" - assert auth_log.event == "exception.client" + _assert_client_auth_log(auth_log) assert auth_log.extra["failure_reason"] == "jwt_invalid" assert auth_log.extra["jwt_algorithm"] == "RS256" assert auth_log.extra["jwt_kid_present"] is True @@ -347,8 +391,7 @@ async def test_unknown_key_id_refreshes_once_and_remains_a_client_warning( assert len(log_capture.records) == 1 auth_log = log_capture.records[0] - assert auth_log.level == "WARNING" - assert auth_log.event == "exception.client" + _assert_client_auth_log(auth_log) assert auth_log.extra["failure_reason"] == "jwt_unknown_key_id" assert auth_log.extra["jwt_algorithm"] == "RS256" assert auth_log.extra["jwt_kid_present"] is True @@ -365,10 +408,11 @@ async def test_unavailable_jwks_is_a_system_error_with_an_unchanged_response( ) -> None: signing_key = _create_rsa_private_key() token = _create_rsa_token(signing_key, key_id="unavailable-key") + jwks_body = b"dashboard-jwks-secret-body" with _capture_auth_logs() as log_capture: with _serve_jwks() as jwks_server: - jwks_server.state.set_raw_response(b"unavailable", status_code=503) + jwks_server.state.set_raw_response(jwks_body, status_code=503) _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) response = await _request_with_token(token) @@ -381,13 +425,13 @@ async def test_unavailable_jwks_is_a_system_error_with_an_unchanged_response( assert len(log_capture.records) == 1 auth_log = log_capture.records[0] - assert auth_log.level == "ERROR" - assert auth_log.event == "exception.system" - assert auth_log.extra["error_category"] == "system" + _assert_jwks_dependency_auth_log(auth_log) + assert auth_log.exception_type == "PyJWKClientConnectionError" assert auth_log.extra["failure_reason"] == "jwks_unavailable" assert auth_log.extra["jwt_kid"] == "unavailable-key" _assert_log_excludes_token(auth_log, token=token) + _assert_log_excludes_jwks_body(auth_log, jwks_body=jwks_body) @pytest.mark.parametrize( @@ -421,13 +465,12 @@ async def test_invalid_jwks_is_a_system_error_with_an_unchanged_response( assert len(log_capture.records) == 1 auth_log = log_capture.records[0] - assert auth_log.level == "ERROR" - assert auth_log.event == "exception.system" - assert auth_log.extra["error_category"] == "system" + _assert_jwks_dependency_auth_log(auth_log) assert auth_log.extra["failure_reason"] == "jwks_invalid" assert auth_log.extra["jwt_kid"] == "invalid-jwks-key" _assert_log_excludes_token(auth_log, token=token) + _assert_log_excludes_jwks_body(auth_log, jwks_body=jwks_body) @pytest.mark.asyncio @@ -481,16 +524,14 @@ async def test_expired_jwt_retains_response_and_logs_client_warning( _assert_unauthenticated_response( response, - expected_message="Token has expired", + expected_message="Authentication required", token=token, ) assert jwks_server.state.request_count == 1 assert len(log_capture.records) == 1 auth_log = log_capture.records[0] - assert auth_log.level == "WARNING" - assert auth_log.event == "exception.client" - assert auth_log.extra["error_category"] == "client" + _assert_client_auth_log(auth_log) assert auth_log.extra["failure_reason"] == "jwt_expired" assert auth_log.extra["jwt_algorithm"] == "RS256" assert auth_log.extra["jwt_kid"] == key_id @@ -516,16 +557,14 @@ async def test_invalid_signature_retains_response_and_logs_client_warning( _assert_unauthenticated_response( response, - expected_message="Invalid token", + expected_message="Authentication required", token=token, ) assert jwks_server.state.request_count == 1 assert len(log_capture.records) == 1 auth_log = log_capture.records[0] - assert auth_log.level == "WARNING" - assert auth_log.event == "exception.client" - assert auth_log.extra["error_category"] == "client" + _assert_client_auth_log(auth_log) assert auth_log.extra["failure_reason"] == "jwt_invalid" assert auth_log.extra["jwt_algorithm"] == "RS256" assert auth_log.extra["jwt_kid"] == key_id @@ -554,16 +593,14 @@ async def test_missing_user_claim_retains_response_and_logs_client_warning( _assert_unauthenticated_response( response, - expected_message="Token missing 'id' claim", + expected_message="Authentication required", token=token, ) assert jwks_server.state.request_count == 1 assert len(log_capture.records) == 1 auth_log = log_capture.records[0] - assert auth_log.level == "WARNING" - assert auth_log.event == "exception.client" - assert auth_log.extra["error_category"] == "client" + _assert_client_auth_log(auth_log) assert auth_log.extra["failure_reason"] == "jwt_invalid" assert auth_log.extra["jwt_algorithm"] == "RS256" assert auth_log.extra["jwt_kid"] == key_id @@ -592,14 +629,13 @@ async def test_non_allowlisted_algorithm_is_not_recorded_as_safe_metadata( _assert_unauthenticated_response( response, - expected_message="Invalid token", + expected_message="Authentication required", token=token, ) assert len(log_capture.records) == 1 auth_log = log_capture.records[0] - assert auth_log.level == "WARNING" - assert auth_log.event == "exception.client" + _assert_client_auth_log(auth_log) assert auth_log.extra["failure_reason"] == "jwt_invalid" assert auth_log.extra["jwt_kid_present"] is True assert auth_log.extra["jwt_kid"] == key_id @@ -607,24 +643,7 @@ async def test_non_allowlisted_algorithm_is_not_recorded_as_safe_metadata( _assert_log_excludes_token(auth_log, token=token) -def test_logfire_exception_callback_keeps_system_category_auth_errors_recorded() -> None: - helper = _FakeLogfireExceptionHelper( - exception=_create_auth_exception( - error_category="system", - exception_context={ - "auth_component": "dashboard_jwt", - "failure_reason": "jwks_unavailable", - }, - ) - ) - - _downgrade_logfire_exception(helper) - - assert helper.level == "error" - assert helper.is_recording_exception is True - - -def test_logfire_exception_callback_downgrades_client_category_auth_errors() -> None: +def test_logfire_exception_callback_downgrades_auth_exceptions_by_status() -> None: helper = _FakeLogfireExceptionHelper(exception=_create_auth_exception()) _downgrade_logfire_exception(helper) diff --git a/packages/shared-python/shared/core/exceptions/domain_exceptions.py b/packages/shared-python/shared/core/exceptions/domain_exceptions.py index 6b8b518be..57ddd79a1 100644 --- a/packages/shared-python/shared/core/exceptions/domain_exceptions.py +++ b/packages/shared-python/shared/core/exceptions/domain_exceptions.py @@ -42,10 +42,9 @@ # User sees: "An internal system error occurred. Please contact support." """ -from collections.abc import Mapping from typing import Any, Dict, List, Optional, TypedDict -from shared.core.exceptions.knowhere_exception import ErrorCategory, KnowhereException +from shared.core.exceptions.knowhere_exception import KnowhereException from shared.core.response.ErrorCode import ErrorCode, SubCode # ============================================================================ @@ -113,16 +112,12 @@ def __init__( self, user_message: str = "Authentication required", internal_message: Optional[str] = None, - error_category: ErrorCategory | None = None, - exception_context: Mapping[str, object] | None = None, ): super().__init__( code=ErrorCode.UNAUTHENTICATED, internal_message=internal_message or user_message, user_message=user_message, details={}, # Empty for security - error_category=error_category, - exception_context=exception_context, ) diff --git a/packages/shared-python/shared/core/exceptions/knowhere_exception.py b/packages/shared-python/shared/core/exceptions/knowhere_exception.py index 22c91fbc3..424431a0b 100644 --- a/packages/shared-python/shared/core/exceptions/knowhere_exception.py +++ b/packages/shared-python/shared/core/exceptions/knowhere_exception.py @@ -59,7 +59,6 @@ raise KnowhereException(code=ErrorCode.INVALID_ARGUMENT, ...) """ -from collections.abc import Mapping from typing import Any, Dict, Literal, Optional from shared.core.response.ErrorCode import ErrorCode, ErrorCodeMapper @@ -67,7 +66,7 @@ # Default messages for auto-sanitization DEFAULT_5XX_USER_MESSAGE = "An internal system error occurred. Please contact support." DEFAULT_4XX_USER_MESSAGE = "Invalid request. Please check your input." -ErrorCategory = Literal["client", "system"] +LogErrorCategory = Literal["client", "system"] class KnowhereException(Exception): @@ -123,8 +122,6 @@ def __init__( details: Optional[Dict[str, Any]] = None, http_status_code: Optional[int] = None, original_exception: Optional[Exception] = None, - error_category: ErrorCategory | None = None, - exception_context: Mapping[str, object] | None = None, ): """ Initialize a KnowhereException. @@ -138,10 +135,6 @@ def __init__( details: Optional structured data to include in response (must be safe). http_status_code: Override HTTP status (auto-derived from code if None). original_exception: The underlying exception being wrapped (for logging). - error_category: Optional telemetry category override. Defaults to the - category derived from the HTTP status. - exception_context: Internal-only structured telemetry fields. These are - included in logs and never returned to clients. """ super().__init__(internal_message) self.code = code @@ -151,11 +144,6 @@ def __init__( http_status_code or ErrorCodeMapper.get_http_status_from_error_code(code) ) self.original_exception = original_exception - default_error_category: ErrorCategory = ( - "system" if self.http_status_code >= 500 else "client" - ) - self.error_category: ErrorCategory = error_category or default_error_category - self.exception_context: Dict[str, object] = dict(exception_context or {}) # ======================================================================= # SECURITY: Auto-sanitize user_message based on HTTP status @@ -229,10 +217,9 @@ def to_log(self) -> Dict[str, Any]: - original_exception: Wrapped exception info """ log_data: Dict[str, Any] = { - **self.exception_context, "error_code": self.code.value, "http_status": self.http_status_code, - "error_category": self.error_category, + "error_category": _get_error_category(self.http_status_code), "exception_class": self.__class__.__name__, "internal_message": self.internal_message, "user_message": self.user_message, @@ -286,9 +273,8 @@ def logging(self, **extra_context): } # Log at appropriate level with appropriate event - if self.error_category == "system": - # System-category errors use ERROR level with a stacktrace even when - # their public HTTP status intentionally remains a 4xx response. + if self.http_status_code >= 500: + # 5xx: ERROR level with stacktrace logger.bind(event=LogEvent.EXCEPTION_SYSTEM.value, **log_data).opt( exception=self ).error(self.internal_message) @@ -336,3 +322,11 @@ def _reconstruct_knowhere_exception(cls, state): obj = cls.__new__(cls) obj.__setstate__(state) return obj + + +def _get_error_category(http_status_code: int) -> LogErrorCategory: + """Return the stable log category derived from the public HTTP status.""" + if http_status_code >= 500: + return "system" + + return "client" diff --git a/packages/shared-python/shared/core/logging.py b/packages/shared-python/shared/core/logging.py index 60b815666..bb246f18e 100644 --- a/packages/shared-python/shared/core/logging.py +++ b/packages/shared-python/shared/core/logging.py @@ -84,7 +84,7 @@ def _is_expected_client_exception(exception: BaseException) -> bool: from shared.core.exceptions.knowhere_exception import KnowhereException if isinstance(exception, KnowhereException): - return exception.error_category == "client" + return 400 <= exception.http_status_code < 500 try: from fastapi import HTTPException as FastAPIHTTPException From fe6b88023b57f56ade23bb1c2ba679c649936751 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 16:59:47 +0800 Subject: [PATCH 12/19] refactor: simplify jwt structure decode options --- .../services/auth/dashboard_jwt_authentication_service.py | 7 ------- 1 file changed, 7 deletions(-) diff --git a/apps/api/app/services/auth/dashboard_jwt_authentication_service.py b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py index 300109e3d..79f9ec989 100644 --- a/apps/api/app/services/auth/dashboard_jwt_authentication_service.py +++ b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py @@ -26,13 +26,6 @@ JWT_ALGORITHMS: tuple[str, ...] = ("HS256", "RS256", "EdDSA") JWT_STRUCTURE_ONLY_DECODE_OPTIONS: Options = { "verify_signature": False, - "verify_exp": False, - "verify_nbf": False, - "verify_iat": False, - "verify_aud": False, - "verify_iss": False, - "verify_sub": False, - "verify_jti": False, } READ_ONLY_PERMISSION: Literal["read_only"] = "read_only" FULL_ACCESS_PERMISSION: Literal["full_access"] = "full_access" From d26b6372ee3aefd08f99b625346927d8c7b33ca3 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 17:58:03 +0800 Subject: [PATCH 13/19] refactor: clarify dashboard jwt auth pipeline --- .../dashboard_jwt_authentication_service.py | 256 +++++++++++------- 1 file changed, 162 insertions(+), 94 deletions(-) diff --git a/apps/api/app/services/auth/dashboard_jwt_authentication_service.py b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py index 79f9ec989..0989460ff 100644 --- a/apps/api/app/services/auth/dashboard_jwt_authentication_service.py +++ b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py @@ -47,6 +47,30 @@ class DashboardJWTIdentity: permission: Permission +@dataclass(frozen=True) +class _DashboardJWTHeader: + algorithm: object + key_id: str + + +@dataclass(frozen=True) +class _DashboardJWTTelemetry: + jwt_kid_present: bool + jwt_algorithm: str | None = None + jwt_kid: str | None = None + + def to_log_data(self) -> dict[str, object]: + log_data: dict[str, object] = { + "auth_component": "dashboard_jwt", + "jwt_kid_present": self.jwt_kid_present, + } + if self.jwt_algorithm is not None: + log_data["jwt_algorithm"] = self.jwt_algorithm + if self.jwt_kid is not None: + log_data["jwt_kid"] = self.jwt_kid + return log_data + + class DashboardJWTAuthenticationService: """Validate Dashboard-issued JWTs through the configured JWKS endpoint.""" @@ -60,55 +84,29 @@ def decode_user_id(self, token: str) -> str: def decode_identity(self, token: str) -> DashboardJWTIdentity: """Decode and validate a JWT, returning the user ID and permission.""" - try: - unverified_header = cast(dict[str, object], jwt.get_unverified_header(token)) - except jwt.InvalidTokenError: - _reject_client_jwt( - failure_reason="jwt_invalid", - telemetry_context=_build_telemetry_context( - algorithm=None, - key_id=None, - ), - ) - - algorithm = unverified_header.get("alg") - key_id_value = unverified_header.get("kid") - key_id = key_id_value if isinstance(key_id_value, str) else None - telemetry_context = _build_telemetry_context( - algorithm=algorithm, - key_id=key_id, + header = _parse_header_or_reject(token) + telemetry_context = _build_telemetry_context(header) + _assert_token_structure_or_reject(token, telemetry_context) + key = self._resolve_verification_key_or_reject( + key_id=header.key_id, + telemetry_context=telemetry_context, ) + payload = _verify_payload_or_reject( + token=token, + key=key, + telemetry_context=telemetry_context, + ) + return _build_identity_or_reject(payload, telemetry_context) - if key_id is None or not key_id.strip(): - _reject_client_jwt( - failure_reason="jwt_missing_key_id", - telemetry_context=telemetry_context, - ) - + def _resolve_verification_key_or_reject( + self, + *, + key_id: str, + telemetry_context: _DashboardJWTTelemetry, + ) -> VerificationKey: + """Resolve a JWT verification key and classify JWKS failures locally.""" try: - self._reject_malformed_token_before_jwks_lookup(token, telemetry_context) key = self._get_verification_key(key_id) - if key is None: - _reject_client_jwt( - failure_reason="jwt_unknown_key_id", - telemetry_context=telemetry_context, - ) - - payload = self._decode_payload(token, key) - user_id = payload.get("id") - if not isinstance(user_id, str) or not user_id: - _reject_client_jwt( - failure_reason="jwt_invalid", - telemetry_context=telemetry_context, - ) - - permission = _normalize_permission(payload.get("permission")) - return DashboardJWTIdentity(user_id=user_id, permission=permission) - except jwt.ExpiredSignatureError: - _reject_client_jwt( - failure_reason="jwt_expired", - telemetry_context=telemetry_context, - ) except PyJWKClientConnectionError as error: _reject_jwks_dependency( failure_reason="jwks_unavailable", @@ -132,48 +130,15 @@ def decode_identity(self, token: str) -> DashboardJWTIdentity: telemetry_context=telemetry_context, original_exception=error, ) - except jwt.InvalidTokenError: - _reject_client_jwt( - failure_reason="jwt_invalid", - telemetry_context=telemetry_context, - ) - - def _decode_payload( - self, - token: str, - key: VerificationKey, - ) -> dict[str, object]: - payload = cast( - dict[str, object], - jwt.decode( - token, - key, - algorithms=list(JWT_ALGORITHMS), - leeway=timedelta(seconds=30), - options={"verify_aud": False}, - ), - ) - return payload - def _reject_malformed_token_before_jwks_lookup( - self, - token: str, - telemetry_context: dict[str, object], - ) -> None: - """Reject structurally invalid JWTs before touching Dashboard JWKS.""" - try: - # This decode only checks token structure; verified claims come from - # _decode_payload after the signing key is resolved. - jwt.decode( - token, - options=JWT_STRUCTURE_ONLY_DECODE_OPTIONS, - ) - except (json.JSONDecodeError, UnicodeDecodeError, jwt.InvalidTokenError): + if key is None: _reject_client_jwt( - failure_reason="jwt_invalid", + failure_reason="jwt_unknown_key_id", telemetry_context=telemetry_context, ) + return key + def _get_verification_key(self, key_id: str) -> VerificationKey | None: """Resolve the JWT verification key from the Dashboard JWKS endpoint.""" jwks_client = self._get_jwks_client() @@ -210,21 +175,124 @@ def _get_jwks_client(self) -> PyJWKClient: return self._jwks_client +def _parse_header_or_reject(token: str) -> _DashboardJWTHeader: + try: + unverified_header = cast(dict[str, object], jwt.get_unverified_header(token)) + except jwt.InvalidTokenError: + _reject_client_jwt( + failure_reason="jwt_invalid", + telemetry_context=_build_telemetry_context_from_values( + algorithm=None, + key_id=None, + ), + ) + + algorithm = unverified_header.get("alg") + key_id_value = unverified_header.get("kid") + key_id = key_id_value if isinstance(key_id_value, str) else None + if key_id is None or not key_id.strip(): + _reject_client_jwt( + failure_reason="jwt_missing_key_id", + telemetry_context=_build_telemetry_context_from_values( + algorithm=algorithm, + key_id=key_id, + ), + ) + + return _DashboardJWTHeader(algorithm=algorithm, key_id=key_id) + + def _build_telemetry_context( + header: _DashboardJWTHeader, +) -> _DashboardJWTTelemetry: + return _build_telemetry_context_from_values( + algorithm=header.algorithm, + key_id=header.key_id, + ) + + +def _build_telemetry_context_from_values( *, algorithm: object, key_id: str | None, -) -> dict[str, object]: +) -> _DashboardJWTTelemetry: is_key_id_present = key_id is not None and bool(key_id.strip()) - context: dict[str, object] = { - "auth_component": "dashboard_jwt", - "jwt_kid_present": is_key_id_present, - } + jwt_algorithm: str | None = None if isinstance(algorithm, str) and algorithm in JWT_ALGORITHMS: - context["jwt_algorithm"] = algorithm + jwt_algorithm = algorithm + jwt_kid: str | None = None if is_key_id_present and key_id is not None: - context["jwt_kid"] = _sanitize_key_id(key_id) - return context + jwt_kid = _sanitize_key_id(key_id) + return _DashboardJWTTelemetry( + jwt_kid_present=is_key_id_present, + jwt_algorithm=jwt_algorithm, + jwt_kid=jwt_kid, + ) + + +def _assert_token_structure_or_reject( + token: str, + telemetry_context: _DashboardJWTTelemetry, +) -> None: + """Reject structurally invalid JWTs before touching Dashboard JWKS.""" + try: + # This decode only checks token structure; verified claims come from + # _verify_payload_or_reject after the signing key is resolved. + jwt.decode( + token, + options=JWT_STRUCTURE_ONLY_DECODE_OPTIONS, + ) + except (json.JSONDecodeError, UnicodeDecodeError, jwt.InvalidTokenError): + _reject_client_jwt( + failure_reason="jwt_invalid", + telemetry_context=telemetry_context, + ) + + +def _verify_payload_or_reject( + *, + token: str, + key: VerificationKey, + telemetry_context: _DashboardJWTTelemetry, +) -> dict[str, object]: + try: + payload = cast( + dict[str, object], + jwt.decode( + token, + key, + algorithms=list(JWT_ALGORITHMS), + leeway=timedelta(seconds=30), + options={"verify_aud": False}, + ), + ) + except jwt.ExpiredSignatureError: + _reject_client_jwt( + failure_reason="jwt_expired", + telemetry_context=telemetry_context, + ) + except (json.JSONDecodeError, UnicodeDecodeError, jwt.InvalidTokenError): + _reject_client_jwt( + failure_reason="jwt_invalid", + telemetry_context=telemetry_context, + ) + + return payload + + +def _build_identity_or_reject( + payload: dict[str, object], + telemetry_context: _DashboardJWTTelemetry, +) -> DashboardJWTIdentity: + user_id = payload.get("id") + if not isinstance(user_id, str) or not user_id: + _reject_client_jwt( + failure_reason="jwt_invalid", + telemetry_context=telemetry_context, + ) + + permission = _normalize_permission(payload.get("permission")) + return DashboardJWTIdentity(user_id=user_id, permission=permission) def _sanitize_key_id(key_id: str) -> str: @@ -235,7 +303,7 @@ def _sanitize_key_id(key_id: str) -> str: def _reject_client_jwt( *, failure_reason: JWTFailureReason, - telemetry_context: dict[str, object], + telemetry_context: _DashboardJWTTelemetry, ) -> NoReturn: _log_dashboard_jwt_auth_failure( failure_reason=failure_reason, @@ -248,7 +316,7 @@ def _reject_client_jwt( def _reject_jwks_dependency( *, failure_reason: JWTFailureReason, - telemetry_context: dict[str, object], + telemetry_context: _DashboardJWTTelemetry, original_exception: Exception, ) -> NoReturn: _log_dashboard_jwt_auth_failure( @@ -263,12 +331,12 @@ def _reject_jwks_dependency( def _log_dashboard_jwt_auth_failure( *, failure_reason: JWTFailureReason, - telemetry_context: dict[str, object], + telemetry_context: _DashboardJWTTelemetry, is_jwks_dependency_failure: bool, original_exception: Exception | None = None, ) -> None: log_data: dict[str, object] = { - **telemetry_context, + **telemetry_context.to_log_data(), "failure_reason": failure_reason, } message = f"Dashboard JWT authentication failed: {failure_reason}" From df7acbc56e6b9b26b1915dbd323dce37b42c84bd Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 19:53:38 +0800 Subject: [PATCH 14/19] feat: add v2 BYOK OpenAI-compatible LLM credentials Allow users to pass text and vision provider configs on v2 job create and retrieval so KNOWHERE can drive parsing and agentic retrieval with caller-supplied OpenAI-compatible keys, models, and endpoints. Co-authored-by: Cursor --- apps/api/app/api/v1/routes/retrieval.py | 30 +++- apps/api/app/api/v2/routes/retrieval.py | 46 ++++- .../connect_builder/summary_builder.py | 5 +- .../document_agent/executor/react_loop.py | 4 +- .../document_agent/planner/planner.py | 4 +- .../structure/page_locate_agent.py | 4 +- .../structure/page_locate_subagent.py | 4 +- .../tools/extract_toc_with_boundaries.py | 5 +- .../document_agent/tools/inspect_pages.py | 4 +- .../tools/propose_shard_plan.py | 4 +- .../document_agent/tools/vlm_toc_extractor.py | 5 +- .../document_ingestion/processing_run.py | 3 + .../formats/fragment/parser.py | 5 +- .../document_parser/formats/image/parser.py | 10 +- .../structure/heading_llm_executor.py | 5 +- .../structure/layout_parser.py | 5 +- .../document_parser/structure/toc_parser.py | 6 +- .../tables/table_frame_parser.py | 5 +- .../services/page_memory/fine_hierarchy.py | 5 +- .../app/services/page_memory/page_assets.py | 8 +- .../app/services/page_memory/page_tagger.py | 9 +- .../test_page_memory_asset_java_contract.py | 2 +- ...est_page_memory_fine_hierarchy_contract.py | 26 +-- ...est_page_memory_node_assembler_contract.py | 2 +- .../shared/models/schemas/job.py | 12 ++ .../shared/models/schemas/job_metadata.py | 48 +++++- .../shared/models/schemas/llm_config.py | 80 +++++++++ .../shared/services/ai/llm_overrides.py | 159 ++++++++++++++++++ .../shared/services/ai/summary/engine.py | 25 ++- .../shared/services/retrieval/app_service.py | 3 + .../services/retrieval/execution/plan.py | 14 ++ .../retrieval/execution/query_request.py | 13 ++ .../shared/services/retrieval/llm_adapter.py | 67 ++++++-- 33 files changed, 546 insertions(+), 81 deletions(-) create mode 100644 packages/shared-python/shared/models/schemas/llm_config.py create mode 100644 packages/shared-python/shared/services/ai/llm_overrides.py diff --git a/apps/api/app/api/v1/routes/retrieval.py b/apps/api/app/api/v1/routes/retrieval.py index 9ed83c8c8..bdba6660e 100644 --- a/apps/api/app/api/v1/routes/retrieval.py +++ b/apps/api/app/api/v1/routes/retrieval.py @@ -2,7 +2,7 @@ from __future__ import annotations -from typing import Literal +from typing import Any, Literal from app.api.dependencies.current_user import with_current_user from app.services.rate_limit.data_structures import CurrentUser @@ -11,6 +11,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from shared.core.database import get_db +from shared.models.schemas.llm_config import LLMConfig from shared.models.schemas.retrieval_namespace import normalize_retrieval_namespace from shared.services.retrieval.app_service import run_retrieval_query from shared.services.retrieval.settings import DEFAULT_TOP_K, VALID_CHUNK_TYPES, normalize_chunk_types @@ -129,12 +130,14 @@ class RetrievalQueryResponse(BaseModel): ) -@router.post("/query", response_model=RetrievalQueryResponse) -async def query_retrieval( +async def execute_retrieval_query( payload: RetrievalQueryRequest, - current_user: CurrentUser = Depends(with_current_user), - db: AsyncSession = Depends(get_db), -): + current_user: CurrentUser, + db: AsyncSession, + *, + llm_config: LLMConfig | None = None, +) -> dict[str, Any]: + """Shared retrieval execution used by v1 and v2 route handlers.""" # Resolve chunk_types: explicit field takes precedence over legacy data_type if payload.chunk_types is not None: resolved_chunk_types = normalize_chunk_types(payload.chunk_types) @@ -166,4 +169,19 @@ async def query_retrieval( threshold=payload.threshold, internal_recall_k=payload.internal_recall_k, use_agentic=payload.use_agentic, + llm_config=llm_config, + ) + + +@router.post("/query", response_model=RetrievalQueryResponse) +async def query_retrieval( + payload: RetrievalQueryRequest, + current_user: CurrentUser = Depends(with_current_user), + db: AsyncSession = Depends(get_db), +): + return await execute_retrieval_query( + payload, + current_user, + db, + llm_config=None, ) diff --git a/apps/api/app/api/v2/routes/retrieval.py b/apps/api/app/api/v2/routes/retrieval.py index e6f188711..0516fb32d 100644 --- a/apps/api/app/api/v2/routes/retrieval.py +++ b/apps/api/app/api/v2/routes/retrieval.py @@ -1,5 +1,47 @@ """Retrieval API v2 routes.""" -from app.api.v1.routes.retrieval import router +from __future__ import annotations -__all__ = ["router"] +from app.api.dependencies.current_user import with_current_user +from app.api.v1.routes.retrieval import ( + RetrievalQueryRequest, + RetrievalQueryResponse, + execute_retrieval_query, +) +from app.services.rate_limit.data_structures import CurrentUser +from fastapi import APIRouter, Depends +from pydantic import Field +from sqlalchemy.ext.asyncio import AsyncSession + +from shared.core.database import get_db +from shared.models.schemas.llm_config import LLMConfig + +router = APIRouter(tags=["Retrieval"]) + + +class RetrievalQueryRequestV2(RetrievalQueryRequest): + """v2 retrieval query request with optional BYOK LLM credentials.""" + + llm_config: LLMConfig | None = Field( + None, + description=( + "Optional bring-your-own-key OpenAI-compatible LLM credentials. " + "Provide text and/or vision provider configs; if only one is set it " + "is used as a unified multimodal model for both. When omitted, " + "server defaults are used." + ), + ) + + +@router.post("/query", response_model=RetrievalQueryResponse) +async def query_retrieval( + payload: RetrievalQueryRequestV2, + current_user: CurrentUser = Depends(with_current_user), + db: AsyncSession = Depends(get_db), +): + return await execute_retrieval_query( + payload, + current_user, + db, + llm_config=payload.llm_config, + ) diff --git a/apps/worker/app/services/connect_builder/summary_builder.py b/apps/worker/app/services/connect_builder/summary_builder.py index f082782b3..960edb836 100644 --- a/apps/worker/app/services/connect_builder/summary_builder.py +++ b/apps/worker/app/services/connect_builder/summary_builder.py @@ -40,7 +40,7 @@ def _llm_summarize(snippets_text: str, node_name: str, max_tokens: int = 100) -> """ try: from shared.services.ai.prompt_service import build_prompt, _detect_text_language - from shared.services.ai.openai_compatible_client_sync import get_openai_client + from shared.services.ai.llm_overrides import get_text_client # Deterministic language lock — see prompt_service._language_directive detected_lang = _detect_text_language(snippets_text) @@ -58,7 +58,8 @@ def _llm_summarize(snippets_text: str, node_name: str, max_tokens: int = 100) -> {"role": "system", "content": "you are a helpful assistant"}, {"role": "user", "content": prompt}, ] - resp = get_openai_client().chat_completion( + client, _ = get_text_client() + resp = client.chat_completion( messages=messages, timeout=60, max_tokens=max_tokens, diff --git a/apps/worker/app/services/document_agent/executor/react_loop.py b/apps/worker/app/services/document_agent/executor/react_loop.py index d391c7647..a890f702f 100644 --- a/apps/worker/app/services/document_agent/executor/react_loop.py +++ b/apps/worker/app/services/document_agent/executor/react_loop.py @@ -235,9 +235,9 @@ def _next_decision(self, round_index: int) -> tuple[ReflexionDecision, ToolResul ) start = time.monotonic() try: - from shared.services.ai.openai_compatible_client_sync import get_openai_client + from shared.services.ai.llm_overrides import get_text_client - client = get_openai_client(model=model) + client, model = get_text_client(requested_model=model) raw, usage = client.chat_completion_with_usage( messages=[{"role": "user", "content": prompt}], model=model, diff --git a/apps/worker/app/services/document_agent/planner/planner.py b/apps/worker/app/services/document_agent/planner/planner.py index b7df8671e..533cc4736 100644 --- a/apps/worker/app/services/document_agent/planner/planner.py +++ b/apps/worker/app/services/document_agent/planner/planner.py @@ -286,9 +286,9 @@ def propose(self) -> tuple[DocumentProfile, ReflexionDecision, ToolResult]: logger.warning("[document_agent] planner png attach failed: {}", exc) try: - from shared.services.ai.openai_compatible_client_sync import get_openai_client + from shared.services.ai.llm_overrides import get_vision_client - client = get_openai_client(model=model) + client, model = get_vision_client(requested_model=model) raw, usage = client.chat_completion_with_usage( messages=cast(Any, [{"role": "user", "content": content_parts}]), model=model, diff --git a/apps/worker/app/services/document_agent/structure/page_locate_agent.py b/apps/worker/app/services/document_agent/structure/page_locate_agent.py index ad033f54f..7a68f3625 100644 --- a/apps/worker/app/services/document_agent/structure/page_locate_agent.py +++ b/apps/worker/app/services/document_agent/structure/page_locate_agent.py @@ -328,9 +328,9 @@ def verify_section_page_choice( start = time.monotonic() try: - from shared.services.ai.openai_compatible_client_sync import get_openai_client + from shared.services.ai.llm_overrides import get_vision_client - client = get_openai_client(model=model) + client, model = get_vision_client(requested_model=model) raw, usage = client.chat_completion_with_usage( messages=cast(Any, [{"role": "user", "content": content_parts}]), model=model, diff --git a/apps/worker/app/services/document_agent/structure/page_locate_subagent.py b/apps/worker/app/services/document_agent/structure/page_locate_subagent.py index a4c9f2dcb..c3826cb2c 100644 --- a/apps/worker/app/services/document_agent/structure/page_locate_subagent.py +++ b/apps/worker/app/services/document_agent/structure/page_locate_subagent.py @@ -299,9 +299,9 @@ def decide_next_action( logger.warning("[page_locate.subagent] planner budget exhausted for title={!r}", title) return None try: - from shared.services.ai.openai_compatible_client_sync import get_openai_client + from shared.services.ai.llm_overrides import get_text_client - client = get_openai_client(model=model) + client, model = get_text_client(requested_model=model) raw, usage = client.chat_completion_with_usage( messages=[{"role": "user", "content": prompt}], model=model, diff --git a/apps/worker/app/services/document_agent/tools/extract_toc_with_boundaries.py b/apps/worker/app/services/document_agent/tools/extract_toc_with_boundaries.py index 33293210a..af2379013 100644 --- a/apps/worker/app/services/document_agent/tools/extract_toc_with_boundaries.py +++ b/apps/worker/app/services/document_agent/tools/extract_toc_with_boundaries.py @@ -70,7 +70,7 @@ def _vlm_confirm_anchors( budget: Any | None = None, ) -> tuple[list[TocAnchorPage], bool, list[TocEvidence]]: """Phase 1: send all anchor PNGs to VLM, ask which are real TOC starts.""" - from shared.services.ai.openai_compatible_client_sync import get_openai_client + from shared.services.ai.llm_overrides import get_vision_client if not anchor_pages: return [], False, [] @@ -125,7 +125,8 @@ def _vlm_confirm_anchors( return [], True, [] try: - client = get_openai_client(model=model) + client, resolved_model = get_vision_client(requested_model=model) + model = resolved_model or model raw, usage = client.chat_completion_with_usage( messages=messages, model=model, diff --git a/apps/worker/app/services/document_agent/tools/inspect_pages.py b/apps/worker/app/services/document_agent/tools/inspect_pages.py index 8881a9b7f..2254dbe47 100644 --- a/apps/worker/app/services/document_agent/tools/inspect_pages.py +++ b/apps/worker/app/services/document_agent/tools/inspect_pages.py @@ -78,9 +78,9 @@ def inspect_pages(ctx: ToolContext, args: dict[str, Any]) -> ToolResult: {"type": "image_url", "image_url": {"url": f"data:image/png;base64,{img_b64}"}} ) try: - from shared.services.ai.openai_compatible_client_sync import get_openai_client + from shared.services.ai.llm_overrides import get_vision_client - client = get_openai_client(model=model) + client, model = get_vision_client(requested_model=model) raw, usage = client.chat_completion_with_usage( messages=cast(Any, [{"role": "user", "content": content_parts}]), model=model, diff --git a/apps/worker/app/services/document_agent/tools/propose_shard_plan.py b/apps/worker/app/services/document_agent/tools/propose_shard_plan.py index 1ca3fb805..714ba9999 100644 --- a/apps/worker/app/services/document_agent/tools/propose_shard_plan.py +++ b/apps/worker/app/services/document_agent/tools/propose_shard_plan.py @@ -738,9 +738,9 @@ def propose_shard_plan(ctx: ToolContext, _args: dict[str, Any]) -> ToolResult: if model and ctx.budget.try_reserve("plan", prompt_tokens_est): try: llm_attempted = True - from shared.services.ai.openai_compatible_client_sync import get_openai_client + from shared.services.ai.llm_overrides import get_text_client - client = get_openai_client(model=model) + client, model = get_text_client(requested_model=model) raw_response, usage = client.chat_completion_with_usage( messages=[{"role": "user", "content": prompt}], model=model, diff --git a/apps/worker/app/services/document_agent/tools/vlm_toc_extractor.py b/apps/worker/app/services/document_agent/tools/vlm_toc_extractor.py index bd5c7c041..9c6c87c95 100644 --- a/apps/worker/app/services/document_agent/tools/vlm_toc_extractor.py +++ b/apps/worker/app/services/document_agent/tools/vlm_toc_extractor.py @@ -162,7 +162,7 @@ def vlm_extract_toc_batch( BatchTocResult with per-page classification and extracted entries. """ from loguru import logger - from shared.services.ai.openai_compatible_client_sync import get_openai_client + from shared.services.ai.llm_overrides import get_vision_client if not page_pngs: return BatchTocResult( @@ -190,7 +190,8 @@ def vlm_extract_toc_batch( ) start = time.monotonic() - client = get_openai_client(model=model) + client, resolved_model = get_vision_client(requested_model=model) + model = resolved_model or model raw, usage = client.chat_completion_with_usage( messages=cast(Any, [{"role": "user", "content": content_parts}]), model=model, diff --git a/apps/worker/app/services/document_ingestion/processing_run.py b/apps/worker/app/services/document_ingestion/processing_run.py index 04975bbf0..4eacc259d 100644 --- a/apps/worker/app/services/document_ingestion/processing_run.py +++ b/apps/worker/app/services/document_ingestion/processing_run.py @@ -34,6 +34,7 @@ ) from shared.core.exceptions.domain_exceptions import ValidationException from shared.models.schemas.job_metadata import JobMetadataHelper +from shared.services.ai.llm_overrides import cleanup_llm_overrides, init_llm_overrides from shared.services.ai.token_tracking import cleanup_token_tracker, init_token_tracker from shared.services.jobs.lifecycle.service import get_sync_job_lifecycle_service from shared.services.redis.distributed_lock import RedisJobLock @@ -87,6 +88,7 @@ def _run_parse_job( lifecycle_service.update_progress(job_id, progress=10, message="Parsing document...") token_usage_dict = init_token_tracker() stage_timing_dict = init_stage_tracker() + init_llm_overrides(JobMetadataHelper.get_llm_config(job_context.job_metadata)) try: prepared_source = prepare_source_file( @@ -180,6 +182,7 @@ def _run_parse_job( "timing_ms": dict(stage_timing_dict), "token_usage": dict(token_usage_dict), } + cleanup_llm_overrides() cleanup_token_tracker() cleanup_stage_tracker() diff --git a/apps/worker/app/services/document_parser/formats/fragment/parser.py b/apps/worker/app/services/document_parser/formats/fragment/parser.py index 1409724bc..059604fb4 100644 --- a/apps/worker/app/services/document_parser/formats/fragment/parser.py +++ b/apps/worker/app/services/document_parser/formats/fragment/parser.py @@ -16,7 +16,7 @@ from openai.types.chat import ChatCompletionMessageParam from app.services.common.file_utils import path_handle -from shared.services.ai.openai_compatible_client_sync import get_openai_client +from shared.services.ai.llm_overrides import get_text_client def generate_fragment_title(content: str, max_tokens: int = 30) -> Optional[str]: @@ -39,7 +39,8 @@ def generate_fragment_title(content: str, max_tokens: int = 30) -> Optional[str] messages: list[ChatCompletionMessageParam] = [ {"role": "user", "content": title_prompt} ] - generated_title = get_openai_client().chat_completion( + client, _ = get_text_client() + generated_title = client.chat_completion( messages=messages, max_tokens=max_tokens, timeout=30, diff --git a/apps/worker/app/services/document_parser/formats/image/parser.py b/apps/worker/app/services/document_parser/formats/image/parser.py index 726e15232..dd7fdd09d 100755 --- a/apps/worker/app/services/document_parser/formats/image/parser.py +++ b/apps/worker/app/services/document_parser/formats/image/parser.py @@ -30,10 +30,8 @@ from app.services.common.file_loading import is_remote, load_file_bytes from app.services.common.file_utils import path_handle from shared.services.ai.summary.engine import summarize, transcribe -from shared.services.ai.openai_compatible_client_sync import ( - OpenAICompatibleClientSync, - get_openai_client, -) +from shared.services.ai.llm_overrides import get_vision_client +from shared.services.ai.openai_compatible_client_sync import OpenAICompatibleClientSync MD_IMAGE_PATTERN = r"!\[[^\]]*?\]\((.*?\.(?:png|jpe?g|gif))\)" g_img_lock = threading.Lock() @@ -61,7 +59,8 @@ def perceptual_hash(data: bytes) -> str: def _get_vision_client() -> OpenAICompatibleClientSync: """Create OpenAI-compatible client for vision models, auto-routing by IMAGE_MODEL name.""" image_model = settings.IMAGE_MODEL or "qwen3.6-flash" - return get_openai_client(model=image_model) + client, _ = get_vision_client(requested_model=image_model) + return client def image_bytes_to_base64(img_data: bytes, ext: str) -> str: @@ -151,6 +150,7 @@ def ask_image( image_model = settings.IMAGE_MODEL_MAX or "qwen3.6-flash" if len(urls_) > 0: + client, image_model = get_vision_client(requested_model=image_model) prompt, temperature, top_p, max_tokens = build_prompt( task=task, texts=title_text, query=query, paras={"max_tokens": max_tokens} ) diff --git a/apps/worker/app/services/document_parser/structure/heading_llm_executor.py b/apps/worker/app/services/document_parser/structure/heading_llm_executor.py index 86a20acde..0af790288 100644 --- a/apps/worker/app/services/document_parser/structure/heading_llm_executor.py +++ b/apps/worker/app/services/document_parser/structure/heading_llm_executor.py @@ -179,7 +179,7 @@ def run_merge_pre_pass( # ── 3. Call LLM directly (bypass df2md — texts is already formatted) ── from shared.services.ai.prompt_service import build_prompt - from shared.services.ai.openai_compatible_client_sync import get_openai_client + from shared.services.ai.llm_overrides import get_text_client from shared.services.ai.response_process_service import eval_response try: @@ -194,7 +194,8 @@ def run_merge_pre_pass( {"role": "user", "content": prompt}, ] with stage_timer("heading.merge_pre_pass_llm", group_count=len(groups), model_name=model_name): - answer = get_openai_client(model=model_name).chat_completion( + client, model_name = get_text_client(requested_model=model_name) + answer = client.chat_completion( messages=messages, model=model_name, max_tokens=max_tokens, diff --git a/apps/worker/app/services/document_parser/structure/layout_parser.py b/apps/worker/app/services/document_parser/structure/layout_parser.py index f99d53297..ea3b8bc83 100755 --- a/apps/worker/app/services/document_parser/structure/layout_parser.py +++ b/apps/worker/app/services/document_parser/structure/layout_parser.py @@ -37,7 +37,7 @@ from shared.services.ai.response_process_service import eval_response # ARQ dependency is removed, use Celery instead -from shared.services.ai.openai_compatible_client_sync import get_openai_client +from shared.services.ai.llm_overrides import get_text_client # ==================== Helper Functions ==================== @@ -230,7 +230,8 @@ def _is_candidate_id(val): candidate_count=n_candidates, max_tokens=max_tokens, ): - answer = get_openai_client(model=model_name).chat_completion( + client, model_name = get_text_client(requested_model=model_name) + answer = client.chat_completion( messages=messages, model=model_name, max_tokens=max_tokens, diff --git a/apps/worker/app/services/document_parser/structure/toc_parser.py b/apps/worker/app/services/document_parser/structure/toc_parser.py index e3a44cca9..884437421 100644 --- a/apps/worker/app/services/document_parser/structure/toc_parser.py +++ b/apps/worker/app/services/document_parser/structure/toc_parser.py @@ -21,7 +21,7 @@ from shared.services.ai.prompt_service import build_prompt from shared.services.ai.response_process_service import eval_response -from shared.services.ai.openai_compatible_client_sync import get_openai_client +from shared.services.ai.llm_overrides import get_text_client # ==================== Markdown TOC Detection Functions ==================== @@ -230,7 +230,9 @@ def llm_judge_toc_range( model_name=model_name, total_candidates=total_candidates, ): - answer = get_openai_client(model=model_name).chat_completion( + client, resolved_model = get_text_client(requested_model=model_name) + model_name = resolved_model or model_name + answer = client.chat_completion( messages=messages, model=model_name, max_tokens=max_tokens, diff --git a/apps/worker/app/services/document_parser/tables/table_frame_parser.py b/apps/worker/app/services/document_parser/tables/table_frame_parser.py index aa9c41678..38aa92b68 100644 --- a/apps/worker/app/services/document_parser/tables/table_frame_parser.py +++ b/apps/worker/app/services/document_parser/tables/table_frame_parser.py @@ -15,7 +15,7 @@ from shared.services.ai.prompt_service import build_prompt from shared.services.ai.response_process_service import eval_response -from shared.services.ai.openai_compatible_client_sync import get_openai_client +from shared.services.ai.llm_overrides import get_text_client from shared.utils.text_utils import remove_duplicates_orderkept @@ -68,7 +68,8 @@ def parse_headers_nonsmart(candidate_frame: pd.DataFrame) -> list[int]: ttl=7200, ) - header_response = get_openai_client().chat_completion( + client, _ = get_text_client() + header_response = client.chat_completion( messages=messages, timeout=60, usage_task="parser.table_detect_headers", diff --git a/apps/worker/app/services/page_memory/fine_hierarchy.py b/apps/worker/app/services/page_memory/fine_hierarchy.py index 361aa072e..0e4466824 100644 --- a/apps/worker/app/services/page_memory/fine_hierarchy.py +++ b/apps/worker/app/services/page_memory/fine_hierarchy.py @@ -19,7 +19,7 @@ from app.services.page_memory.page_tagger import PageTagResult from app.services.page_memory.skeleton_extractor import SectionSkeleton from app.services.page_memory._utils import page_scope_info -from shared.services.ai.openai_compatible_client_sync import get_openai_client +from shared.services.ai.llm_overrides import get_text_client from shared.services.ai.prompt_service import build_prompt from shared.services.ai.response_process_service import eval_response @@ -226,7 +226,8 @@ def _run_hierarchy_on_candidates( os.environ.get("NORMOL_MODEL"), ) ) - answer = get_openai_client(model=resolved_model).chat_completion( + client, resolved_model = get_text_client(requested_model=resolved_model) + answer = client.chat_completion( messages=[ {"role": "system", "content": "you are a document structure expert"}, {"role": "user", "content": prompt}, diff --git a/apps/worker/app/services/page_memory/page_assets.py b/apps/worker/app/services/page_memory/page_assets.py index 232bfc83e..7e067a71c 100644 --- a/apps/worker/app/services/page_memory/page_assets.py +++ b/apps/worker/app/services/page_memory/page_assets.py @@ -109,9 +109,9 @@ def detect_page_assets( ) try: - from shared.services.ai.openai_compatible_client_sync import get_openai_client + from shared.services.ai.llm_overrides import get_vision_client - client = get_openai_client(model=model_name) + client, model_name = get_vision_client(requested_model=model_name) raw_response, usage = client.chat_completion_with_usage( messages=cast( Any, @@ -829,9 +829,9 @@ def _llm_judge_table_continuity( ) try: - from shared.services.ai.openai_compatible_client_sync import get_openai_client + from shared.services.ai.llm_overrides import get_text_client - client = get_openai_client(model=model_name) + client, model_name = get_text_client(requested_model=model_name) raw_response, usage = client.chat_completion_with_usage( messages=[{"role": "user", "content": prompt}], model=model_name, diff --git a/apps/worker/app/services/page_memory/page_tagger.py b/apps/worker/app/services/page_memory/page_tagger.py index 4fe72ac23..4476e6a69 100644 --- a/apps/worker/app/services/page_memory/page_tagger.py +++ b/apps/worker/app/services/page_memory/page_tagger.py @@ -173,9 +173,9 @@ def _tag_text_only( ) try: - from shared.services.ai.openai_compatible_client_sync import get_openai_client + from shared.services.ai.llm_overrides import get_text_client - client = get_openai_client(model=model) + client, model = get_text_client(requested_model=model) raw_response, usage = client.chat_completion_with_usage( messages=cast(Any, [{"role": "user", "content": prompt}]), model=model, @@ -388,9 +388,10 @@ def _tag_vlm_titles( }, ] - from shared.services.ai.openai_compatible_client_sync import get_openai_client + from shared.services.ai.llm_overrides import get_vision_client - client = get_openai_client(model=model) + client, resolved_model = get_vision_client(requested_model=model) + model = resolved_model or model for attempt in range(_MAX_JSON_RETRIES + 1): try: diff --git a/apps/worker/tests/contract/test_page_memory_asset_java_contract.py b/apps/worker/tests/contract/test_page_memory_asset_java_contract.py index 7725d396a..78f13946b 100644 --- a/apps/worker/tests/contract/test_page_memory_asset_java_contract.py +++ b/apps/worker/tests/contract/test_page_memory_asset_java_contract.py @@ -46,7 +46,7 @@ def chat_completion_with_usage(self, **kwargs): import shared.services.ai.openai_compatible_client_sync as client_mod monkeypatch.setattr( - client_mod, "get_openai_client", lambda model=None: _FakeClient() + client_mod, "get_openai_client", lambda model=None, **_kwargs: _FakeClient() ) image_path = tmp_path / "page.png" diff --git a/apps/worker/tests/contract/test_page_memory_fine_hierarchy_contract.py b/apps/worker/tests/contract/test_page_memory_fine_hierarchy_contract.py index c46eac986..4eac1d249 100644 --- a/apps/worker/tests/contract/test_page_memory_fine_hierarchy_contract.py +++ b/apps/worker/tests/contract/test_page_memory_fine_hierarchy_contract.py @@ -98,8 +98,11 @@ def test_refine_fat_leaf_skeletons_excludes_next_section_start_when_unordered( monkeypatch.setattr( fine_hierarchy, - "get_openai_client", - lambda model=None: _FakeClient([{"id": 1, "level": 1}]), + "get_text_client", + lambda requested_model=None: ( + _FakeClient([{"id": 1, "level": 1}]), + requested_model, + ), ) refined = fine_hierarchy.refine_fat_leaf_skeletons( @@ -152,14 +155,17 @@ def test_refine_fat_leaf_skeletons_uses_page_memory_prompt_without_demoting_sibl monkeypatch.setattr( fine_hierarchy, - "get_openai_client", - lambda model=None: _FakeClient( - [ - {"id": 1, "level": 1}, - {"id": 2, "level": 1}, - {"id": 3, "level": 2}, - {"id": 4, "level": 1}, - ] + "get_text_client", + lambda requested_model=None: ( + _FakeClient( + [ + {"id": 1, "level": 1}, + {"id": 2, "level": 1}, + {"id": 3, "level": 2}, + {"id": 4, "level": 1}, + ] + ), + requested_model, ), ) diff --git a/apps/worker/tests/contract/test_page_memory_node_assembler_contract.py b/apps/worker/tests/contract/test_page_memory_node_assembler_contract.py index 21304323e..a5e2edb26 100644 --- a/apps/worker/tests/contract/test_page_memory_node_assembler_contract.py +++ b/apps/worker/tests/contract/test_page_memory_node_assembler_contract.py @@ -410,7 +410,7 @@ def chat_completion_with_usage(self, **kwargs): import shared.services.ai.openai_compatible_client_sync as client_mod monkeypatch.setattr( - client_mod, "get_openai_client", lambda model=None: _FakeClient() + client_mod, "get_openai_client", lambda model=None, **_kwargs: _FakeClient() ) img = tmp_path / "page-231.png" diff --git a/packages/shared-python/shared/models/schemas/job.py b/packages/shared-python/shared/models/schemas/job.py index 4326ba439..c772bf680 100644 --- a/packages/shared-python/shared/models/schemas/job.py +++ b/packages/shared-python/shared/models/schemas/job.py @@ -5,6 +5,8 @@ from pydantic import BaseModel, ConfigDict, Field +from shared.models.schemas.llm_config import LLMConfig + class WebhookConfig(BaseModel): """Webhook configuration.""" @@ -85,6 +87,16 @@ class JobCreate(JobCreateBase): class JobCreateV2(JobCreateBase): """Public v2 request payload for creating a job.""" + llm_config: Optional[LLMConfig] = Field( + None, + description=( + "Optional bring-your-own-key OpenAI-compatible LLM credentials. " + "Provide text and/or vision provider configs; if only one is set it " + "is used as a unified multimodal model for both. When omitted, " + "server defaults are used." + ), + ) + class JobResponse(BaseModel): """Job-creation response.""" diff --git a/packages/shared-python/shared/models/schemas/job_metadata.py b/packages/shared-python/shared/models/schemas/job_metadata.py index 30934dc5f..9e25a2d7a 100644 --- a/packages/shared-python/shared/models/schemas/job_metadata.py +++ b/packages/shared-python/shared/models/schemas/job_metadata.py @@ -4,8 +4,10 @@ from pydantic import BaseModel, ConfigDict, Field +from shared.models.schemas.llm_config import LLMConfig, parse_llm_config from shared.models.schemas.page_memory_config import PageMemoryConfig from shared.models.schemas.retrieval_namespace import normalize_retrieval_namespace +from shared.utils.security_utils import mask_api_key class JobMetadataBase(BaseModel): @@ -18,6 +20,9 @@ class JobMetadataBase(BaseModel): parsing_params: Optional[Dict[str, Any]] = Field( None, description="Parsing parameters" ) + llm_config: Optional[Dict[str, Any]] = Field( + None, description="BYOK OpenAI-compatible LLM credentials (v2)" + ) data_id: Optional[str] = Field(None, description="User-defined ID") webhook: Optional[Dict[str, Any]] = Field(None, description="Webhook configuration") document_metadata: Optional[Dict[str, Any]] = Field( @@ -67,6 +72,13 @@ def create_from_request( resolved_page_memory_config = page_memory_config.to_dict() else: resolved_page_memory_config = page_memory_config + # v2-only field; getattr keeps v1 JobCreate (no llm_config) safe. + raw_llm_config = getattr(request, "llm_config", None) + llm_config_payload: Dict[str, Any] | None = None + if isinstance(raw_llm_config, LLMConfig): + llm_config_payload = raw_llm_config.model_dump() + elif isinstance(raw_llm_config, dict): + llm_config_payload = dict(raw_llm_config) metadata = { "original_request": _dump_public_request(request), "api_version": api_version, @@ -81,6 +93,8 @@ def create_from_request( "data_id": request.data_id, "webhook": request.webhook.model_dump() if request.webhook else None, } + if llm_config_payload is not None: + metadata["llm_config"] = llm_config_payload if resolved_page_memory_config is not None: metadata["page_memory_config"] = resolved_page_memory_config metadata.update(kwargs) @@ -239,8 +253,40 @@ def get_webhook(metadata: Optional[Dict[str, Any]]) -> Optional[Dict[str, Any]]: """Return the webhook configuration from metadata.""" return JobMetadataHelper.get_field(metadata, "webhook") + @staticmethod + def get_llm_config(metadata: Optional[Dict[str, Any]]) -> LLMConfig | None: + """Return the BYOK LLM config stored in metadata, if any.""" + raw = JobMetadataHelper.get_field(metadata, "llm_config", None) + return parse_llm_config(raw) + + +def _mask_llm_config_in_request(payload: Dict[str, Any]) -> Dict[str, Any]: + """Redact api_key values inside llm_config for public request snapshots.""" + llm_config = payload.get("llm_config") + if not isinstance(llm_config, dict): + return payload + + masked = dict(payload) + masked_llm: Dict[str, Any] = {} + for slot in ("text", "vision"): + provider = llm_config.get(slot) + if isinstance(provider, dict): + provider_copy = dict(provider) + if "api_key" in provider_copy: + provider_copy["api_key"] = mask_api_key( + provider_copy.get("api_key") + if isinstance(provider_copy.get("api_key"), str) + else None + ) + masked_llm[slot] = provider_copy + elif provider is not None: + masked_llm[slot] = provider + masked["llm_config"] = masked_llm + return masked + def _dump_public_request(request) -> Dict[str, Any]: """Dump declared public request fields without hidden compatibility extras.""" extra_fields = getattr(request, "model_extra", None) or {} - return request.model_dump(exclude=set(extra_fields)) + payload = request.model_dump(exclude=set(extra_fields)) + return _mask_llm_config_in_request(payload) diff --git a/packages/shared-python/shared/models/schemas/llm_config.py b/packages/shared-python/shared/models/schemas/llm_config.py new file mode 100644 index 000000000..4b71cfe4d --- /dev/null +++ b/packages/shared-python/shared/models/schemas/llm_config.py @@ -0,0 +1,80 @@ +"""Bring-your-own-key (BYOK) OpenAI-compatible LLM credentials.""" + +from __future__ import annotations + +from typing import Any, Optional + +from pydantic import BaseModel, Field, model_validator + + +class LLMProviderConfig(BaseModel): + """Credentials for one OpenAI-compatible provider endpoint.""" + + api_key: str = Field(..., min_length=1, description="Provider API key") + model: str = Field(..., min_length=1, description="Model identifier") + base_url: str = Field( + ..., + min_length=1, + description="OpenAI-compatible base URL (e.g. https://api.openai.com/v1)", + ) + + +class LLMConfig(BaseModel): + """Optional text + vision provider configs for BYOK. + + Semantics: + - both set -> text for text tasks, vision for VLM tasks + - exactly one set -> that config is used as a unified multimodal model + for both text and vision tasks + - neither set -> invalid when the object itself is present + """ + + text: Optional[LLMProviderConfig] = Field( + None, description="Text / planning LLM credentials" + ) + vision: Optional[LLMProviderConfig] = Field( + None, description="Vision / VLM credentials" + ) + + @model_validator(mode="after") + def _require_at_least_one_provider(self) -> "LLMConfig": + if self.text is None and self.vision is None: + raise ValueError("llm_config requires at least one of text or vision") + return self + + def text_effective(self) -> LLMProviderConfig | None: + """Resolve the config used for text tasks (unified-multimodal aware).""" + return self.text if self.text is not None else self.vision + + def vision_effective(self) -> LLMProviderConfig | None: + """Resolve the config used for vision/VLM tasks (unified-multimodal aware).""" + return self.vision if self.vision is not None else self.text + + def masked_dump(self) -> dict[str, Any]: + """Serialize with api_key values redacted for snapshots / responses.""" + from shared.utils.security_utils import mask_api_key + + def _mask_provider(provider: LLMProviderConfig | None) -> dict[str, Any] | None: + if provider is None: + return None + return { + "api_key": mask_api_key(provider.api_key), + "model": provider.model, + "base_url": provider.base_url, + } + + return { + "text": _mask_provider(self.text), + "vision": _mask_provider(self.vision), + } + + +def parse_llm_config(value: Any) -> LLMConfig | None: + """Parse a raw mapping / LLMConfig into a validated LLMConfig, or None.""" + if value is None: + return None + if isinstance(value, LLMConfig): + return value + if isinstance(value, dict): + return LLMConfig.model_validate(value) + return None diff --git a/packages/shared-python/shared/services/ai/llm_overrides.py b/packages/shared-python/shared/services/ai/llm_overrides.py new file mode 100644 index 000000000..7c65fa4be --- /dev/null +++ b/packages/shared-python/shared/services/ai/llm_overrides.py @@ -0,0 +1,159 @@ +"""Request-scoped BYOK LLM credential overrides. + +In the gevent worker, child greenlets (GeventPool.spawn) do NOT inherit +``ContextVar`` or ``threading.local`` from the parent. We therefore use +a module-level dict keyed by the *root* greenlet id of the current parse +task — the same pattern as ``token_tracking``. + +For the asyncio retrieval path a ``ContextVar`` is used instead, which +propagates naturally across ``await`` and ``asyncio.to_thread``. + +``get_current_llm_overrides()`` checks the greenlet dict first, then the +ContextVar, so both runtimes share one lookup API. +""" + +from __future__ import annotations + +import threading +from contextvars import ContextVar, Token +from typing import Any, Optional + +from shared.models.schemas.llm_config import LLMConfig, LLMProviderConfig, parse_llm_config + +# Greenlet-root keyed store (worker / gevent). +_overrides: dict[int, LLMConfig] = {} +_root_ids: dict[int, int] = {} +_lock = threading.Lock() + +# Async ContextVar store (API retrieval). +_async_overrides: ContextVar[LLMConfig | None] = ContextVar( + "llm_overrides_async", + default=None, +) + +ResolvedCredentials = tuple[str | None, str | None, str | None] +# (model, api_key, api_url) + + +def _current_greenlet_id() -> int: + try: + import gevent + + return id(gevent.getcurrent()) + except ImportError: + return threading.get_ident() + + +def _find_root_id() -> int | None: + """Walk up the greenlet parent chain to find a registered root id.""" + gid = _current_greenlet_id() + if gid in _overrides: + return gid + if gid in _root_ids: + return _root_ids[gid] + try: + import gevent + + g = gevent.getcurrent() + while g is not None: + pid = id(g) + if pid in _overrides: + _root_ids[gid] = pid + return pid + g = getattr(g, "parent", None) + except ImportError: + pass + return None + + +def init_llm_overrides(config: LLMConfig | dict[str, Any] | None) -> LLMConfig | None: + """Register BYOK overrides for the current parse-task root greenlet. + + Must be called from the root greenlet of the task (same scope as + ``init_token_tracker``). Returns the parsed config, or None when + no override is active. + """ + parsed = parse_llm_config(config) + if parsed is None: + return None + gid = _current_greenlet_id() + with _lock: + _overrides[gid] = parsed + return parsed + + +def cleanup_llm_overrides() -> None: + """Remove the override for the current greenlet. Call after parsing.""" + gid = _current_greenlet_id() + with _lock: + _overrides.pop(gid, None) + stale = [k for k, v in _root_ids.items() if v == gid] + for k in stale: + del _root_ids[k] + + +def set_llm_overrides_async( + config: LLMConfig | dict[str, Any] | None, +) -> Token[LLMConfig | None]: + """Set BYOK overrides for the current asyncio task. Returns a reset token.""" + parsed = parse_llm_config(config) + return _async_overrides.set(parsed) + + +def reset_llm_overrides_async(token: Token[LLMConfig | None]) -> None: + """Reset the async override ContextVar using the token from set_*.""" + _async_overrides.reset(token) + + +def get_current_llm_overrides() -> LLMConfig | None: + """Return the active BYOK config for this task, if any. + + Checks the gevent greenlet-root store first, then the async ContextVar. + """ + root = _find_root_id() + if root is not None: + return _overrides.get(root) + return _async_overrides.get() + + +def _resolve_provider( + provider: LLMProviderConfig | None, + requested_model: str | None, +) -> ResolvedCredentials: + if provider is None: + return (requested_model, None, None) + return (provider.model, provider.api_key, provider.base_url) + + +def resolve_text(requested_model: str | None = None) -> ResolvedCredentials: + """Resolve (model, api_key, api_url) for a text LLM call.""" + overrides = get_current_llm_overrides() + if overrides is None: + return (requested_model, None, None) + return _resolve_provider(overrides.text_effective(), requested_model) + + +def resolve_vision(requested_model: str | None = None) -> ResolvedCredentials: + """Resolve (model, api_key, api_url) for a vision/VLM call.""" + overrides = get_current_llm_overrides() + if overrides is None: + return (requested_model, None, None) + return _resolve_provider(overrides.vision_effective(), requested_model) + + +def get_text_client(requested_model: Optional[str] = None): + """Return ``(client, effective_model)`` for a text LLM call.""" + from shared.services.ai.openai_compatible_client_sync import get_openai_client + + model, api_key, api_url = resolve_text(requested_model) + client = get_openai_client(model=model, api_key=api_key, api_url=api_url) + return client, model + + +def get_vision_client(requested_model: Optional[str] = None): + """Return ``(client, effective_model)`` for a vision/VLM call.""" + from shared.services.ai.openai_compatible_client_sync import get_openai_client + + model, api_key, api_url = resolve_vision(requested_model) + client = get_openai_client(model=model, api_key=api_key, api_url=api_url) + return client, model diff --git a/packages/shared-python/shared/services/ai/summary/engine.py b/packages/shared-python/shared/services/ai/summary/engine.py index eaedd05c2..141520f6f 100644 --- a/packages/shared-python/shared/services/ai/summary/engine.py +++ b/packages/shared-python/shared/services/ai/summary/engine.py @@ -108,6 +108,7 @@ def _call_llm( budget: Any | None, budget_pool: str, budget_stage: str | None, + channel: Literal["text", "vision"] = "text", ) -> Any | None: """One text-or-vision call with budget accounting and a single JSON retry. @@ -115,7 +116,11 @@ def _call_llm( failure / exhausted budget. Budget is reserved before the call, committed on success, refunded on failure — matching the prior per-caller bookkeeping but in one place. + + ``channel`` selects BYOK text vs vision credentials when overrides are active. """ + from shared.services.ai.llm_overrides import resolve_text, resolve_vision + content_parts: list[dict[str, Any]] = [{"type": "text", "text": prompt}] for path in image_paths: img_b64 = _read_image_b64(path) @@ -142,12 +147,23 @@ def _call_llm( if expect_json: api_kwargs["response_format"] = {"type": "json_object"} - client = _client_mod.get_openai_client(model=model) + resolve = resolve_vision if channel == "vision" else resolve_text + effective_model, api_key, api_url = resolve(model) + if not effective_model: + if budget is not None: + budget.refund(budget_pool, est=est, stage=budget_stage) + return None + + client = _client_mod.get_openai_client( + model=effective_model, + api_key=api_key, + api_url=api_url, + ) for attempt in range(_MAX_JSON_RETRIES + 1): try: raw, usage = client.chat_completion_with_usage( messages=cast(Any, [{"role": "user", "content": content_parts}]), - model=model, + model=effective_model, temperature=temperature, top_p=top_p, max_tokens=max_tokens, @@ -337,6 +353,7 @@ def _summarize_body( budget=budget, budget_pool=budget_pool, budget_stage=budget_stage, + channel="vision", ) else: # Text summary: the shared summary-full prompt with deterministic lang lock. @@ -366,6 +383,7 @@ def _summarize_body( budget=budget, budget_pool="plan", budget_stage=None, + channel="text", ) if not isinstance(parsed, dict): @@ -416,6 +434,7 @@ def _summarize_asset( budget=budget, budget_pool="plan", budget_stage=None, + channel="text", ) if isinstance(raw, str) and raw.strip(): return _parse_linesplit_asset(raw, asset_title_hint) @@ -444,6 +463,7 @@ def _summarize_asset( budget=budget, budget_pool=budget_pool, budget_stage=budget_stage, + channel="vision", ) if isinstance(parsed, dict): return AssetSummary( @@ -492,6 +512,7 @@ def transcribe( budget=budget, budget_pool=budget_pool, budget_stage=budget_stage, + channel="vision", ) if isinstance(parsed, dict): return str(parsed.get("text", "")).strip() diff --git a/packages/shared-python/shared/services/retrieval/app_service.py b/packages/shared-python/shared/services/retrieval/app_service.py index 1e061510a..83ce65b78 100644 --- a/packages/shared-python/shared/services/retrieval/app_service.py +++ b/packages/shared-python/shared/services/retrieval/app_service.py @@ -4,6 +4,7 @@ from sqlalchemy.ext.asyncio import AsyncSession +from shared.models.schemas.llm_config import LLMConfig from shared.models.schemas.retrieval_namespace import normalize_retrieval_namespace from shared.services.retrieval.execution.plan import ( run_retrieval_query as execute_retrieval_query, @@ -31,6 +32,7 @@ async def run_retrieval_query( threshold: float = 0.0, internal_recall_k: int | None = None, use_agentic: bool | None = None, + llm_config: LLMConfig | None = None, ) -> dict[str, Any]: return await execute_retrieval_query( db=db, @@ -49,4 +51,5 @@ async def run_retrieval_query( threshold=threshold, internal_recall_k=internal_recall_k, use_agentic=use_agentic, + llm_config=llm_config, ) diff --git a/packages/shared-python/shared/services/retrieval/execution/plan.py b/packages/shared-python/shared/services/retrieval/execution/plan.py index 99e012e07..ed9f9f591 100644 --- a/packages/shared-python/shared/services/retrieval/execution/plan.py +++ b/packages/shared-python/shared/services/retrieval/execution/plan.py @@ -6,6 +6,11 @@ from loguru import logger from sqlalchemy.ext.asyncio import AsyncSession +from shared.models.schemas.llm_config import LLMConfig +from shared.services.ai.llm_overrides import ( + reset_llm_overrides_async, + set_llm_overrides_async, +) from shared.services.retrieval.cache_service import ( get_cached_retrieval_query_result, set_cached_retrieval_query_result, @@ -38,6 +43,7 @@ async def run_retrieval_query( threshold: float = 0.0, internal_recall_k: int | None = None, use_agentic: bool | None = None, + llm_config: LLMConfig | None = None, ) -> dict[str, Any]: """Run retrieval through the plan module.""" return await RetrievalExecutionPlan( @@ -58,6 +64,7 @@ async def run_retrieval_query( threshold=threshold, internal_recall_k=internal_recall_k, use_agentic=use_agentic, + llm_config=llm_config, ) ).execute() @@ -68,7 +75,14 @@ def __init__(self, request: RetrievalQuery) -> None: async def execute(self) -> dict[str, Any]: request = self.request + override_token = set_llm_overrides_async(request.llm_config) + + try: + return await self._execute_with_overrides(request) + finally: + reset_llm_overrides_async(override_token) + async def _execute_with_overrides(self, request: RetrievalQuery) -> dict[str, Any]: # TODO(intent-step): Insert Intent Understanding step here. # Before any retrieval runs, parse `request.query` with LLM + # KG overview + section tree to extract structured navigation diff --git a/packages/shared-python/shared/services/retrieval/execution/query_request.py b/packages/shared-python/shared/services/retrieval/execution/query_request.py index a7094b888..c226a9f6d 100644 --- a/packages/shared-python/shared/services/retrieval/execution/query_request.py +++ b/packages/shared-python/shared/services/retrieval/execution/query_request.py @@ -5,6 +5,7 @@ from sqlalchemy.ext.asyncio import AsyncSession +from shared.models.schemas.llm_config import LLMConfig from shared.models.schemas.retrieval_namespace import normalize_retrieval_namespace from shared.services.retrieval.execution.route_types import RetrievalRouteContext from shared.services.retrieval.settings import ( @@ -30,6 +31,7 @@ class RetrievalQuery: threshold: float = 0.0 internal_recall_k: int | None = None use_agentic: bool | None = None + llm_config: LLMConfig | None = None @classmethod def from_parameters( @@ -51,6 +53,7 @@ def from_parameters( threshold: float = 0.0, internal_recall_k: int | None = None, use_agentic: bool | None = None, + llm_config: LLMConfig | None = None, ) -> "RetrievalQuery": return cls( db=db, @@ -69,9 +72,17 @@ def from_parameters( threshold=threshold, internal_recall_k=internal_recall_k, use_agentic=use_agentic, + llm_config=llm_config, ) def build_cache_extra(self) -> dict[str, Any]: + text_model: str | None = None + vision_model: str | None = None + if self.llm_config is not None: + text_provider = self.llm_config.text_effective() + vision_provider = self.llm_config.vision_effective() + text_model = text_provider.model if text_provider is not None else None + vision_model = vision_provider.model if vision_provider is not None else None return { "chunk_types": sorted(self.chunk_types) if self.chunk_types else None, "signal_paths": self.signal_paths, @@ -83,6 +94,8 @@ def build_cache_extra(self) -> dict[str, Any]: "internal_recall_k": self.internal_recall_k, "use_agentic": self.use_agentic, "decomposition_enabled": True, + "llm_text_model": text_model, + "llm_vision_model": vision_model, } def resolve_allowed_chunk_types(self) -> set[str] | None: diff --git a/packages/shared-python/shared/services/retrieval/llm_adapter.py b/packages/shared-python/shared/services/retrieval/llm_adapter.py index 9ee344014..8c8cb627d 100644 --- a/packages/shared-python/shared/services/retrieval/llm_adapter.py +++ b/packages/shared-python/shared/services/retrieval/llm_adapter.py @@ -13,6 +13,7 @@ from loguru import logger from shared.core.config import settings +from shared.services.ai.llm_overrides import get_current_llm_overrides # LLMFn accepts either a plain string or a list of ChatCompletionMessageParam LLMFnInput = Union[str, Sequence[dict[str, Any]]] @@ -29,6 +30,8 @@ def _has_llm_credentials() -> bool: """Check whether at least one LLM provider is configured.""" + if get_current_llm_overrides() is not None: + return True if getattr(settings, 'LLM_MOCK_ENABLED', False): return True if getattr(settings, 'DS_KEY', ''): @@ -44,6 +47,11 @@ def _has_llm_credentials() -> bool: def _resolve_default_model() -> str: """Pick a model name that matches the configured LLM provider.""" + overrides = get_current_llm_overrides() + if overrides is not None: + provider = overrides.text_effective() + if provider is not None: + return provider.model if getattr(settings, 'DS_KEY', ''): return 'deepseek-v4-flash' if getattr(settings, 'ALI_API_KEYS', ''): @@ -56,6 +64,11 @@ def _resolve_default_model() -> str: def _resolve_planner_model(*, thinking: bool) -> str: + overrides = get_current_llm_overrides() + if overrides is not None: + provider = overrides.text_effective() + if provider is not None: + return provider.model configured = getattr(settings, 'RETRIEVAL_PLANNER_MODEL', '') or '' if configured: return configured @@ -70,6 +83,29 @@ def _resolve_planner_model(*, thinking: bool) -> str: return getattr(settings, 'NORMOL_MODEL', None) or 'deepseek-v4-flash' +def _resolve_vlm_model(model: str | None = None) -> str: + overrides = get_current_llm_overrides() + if overrides is not None: + provider = overrides.vision_effective() + if provider is not None: + return provider.model + return model or getattr(settings, 'IMAGE_MODEL', '') or 'qwen3.6-flash' + + +def _build_client_for_channel(*, channel: str, model: str): + """Build an OpenAI-compatible client, honoring active BYOK overrides.""" + from shared.services.ai.openai_compatible_client_sync import get_openai_client + from shared.services.ai.llm_overrides import resolve_text, resolve_vision + + resolve = resolve_vision if channel == 'vision' else resolve_text + effective_model, api_key, api_url = resolve(model) + return get_openai_client( + model=effective_model, + api_key=api_key, + api_url=api_url, + ), effective_model + + def create_retrieval_llm_fn( *, model: str | None = None, @@ -108,9 +144,10 @@ def create_retrieval_llm_fn( ) async def llm_fn(prompt: LLMFnInput) -> str: - from shared.services.ai.openai_compatible_client_sync import get_openai_client - - client = get_openai_client(model=effective_model) + client, resolved_model = _build_client_for_channel( + channel='text', + model=effective_model, + ) current_llm_usage.set(None) kwargs: dict[str, Any] = {} @@ -124,7 +161,7 @@ async def llm_fn(prompt: LLMFnInput) -> str: result, usage = await asyncio.to_thread( client.chat_completion_with_usage, cast(Any, prompt), - model=effective_model, + model=resolved_model, temperature=effective_temperature, max_tokens=effective_max_tokens, **kwargs, @@ -149,14 +186,15 @@ def create_retrieval_planner_fn( effective_model = model or _resolve_planner_model(thinking=thinking) async def llm_fn(prompt: LLMFnInput) -> str: - from shared.services.ai.openai_compatible_client_sync import get_openai_client - - client = get_openai_client(model=effective_model) + client, resolved_model = _build_client_for_channel( + channel='text', + model=effective_model, + ) current_llm_usage.set(None) result, usage = await asyncio.to_thread( client.chat_completion_with_usage, cast(Any, prompt), - model=effective_model, + model=resolved_model, temperature=0.0, max_tokens=max_tokens, ) @@ -181,23 +219,22 @@ def create_retrieval_vlm_fn( ``create_retrieval_llm_fn`` — callers pass either a plain string or a list of ChatCompletionMessageParam (including image_url parts). """ - from shared.core.config import settings - - effective_model = model or getattr(settings, 'IMAGE_MODEL', '') or 'qwen3.6-flash' + effective_model = _resolve_vlm_model(model) if not _has_llm_credentials(): logger.debug('retrieval: no LLM credentials for VLM, image-aware answering disabled') return None async def vlm_fn(prompt: LLMFnInput) -> str: - from shared.services.ai.openai_compatible_client_sync import get_openai_client - - client = get_openai_client(model=effective_model) + client, resolved_model = _build_client_for_channel( + channel='vision', + model=effective_model, + ) current_llm_usage.set(None) result, usage = await asyncio.to_thread( client.chat_completion_with_usage, cast(Any, prompt), - model=effective_model, + model=resolved_model, temperature=temperature, max_tokens=max_tokens, ) From e3316ee5c424f406a2e5adb3d1685747dc2ab83d Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 20:08:05 +0800 Subject: [PATCH 15/19] fix: avoid empty except in llm_overrides greenlet walk Return early when gevent is unavailable so CodeQL no longer flags a bare ImportError pass in the BYOK override root lookup. Co-authored-by: Cursor --- .../shared/services/ai/llm_overrides.py | 19 ++++++++++--------- 1 file changed, 10 insertions(+), 9 deletions(-) diff --git a/packages/shared-python/shared/services/ai/llm_overrides.py b/packages/shared-python/shared/services/ai/llm_overrides.py index 7c65fa4be..bc5671ecd 100644 --- a/packages/shared-python/shared/services/ai/llm_overrides.py +++ b/packages/shared-python/shared/services/ai/llm_overrides.py @@ -53,16 +53,17 @@ def _find_root_id() -> int | None: return _root_ids[gid] try: import gevent - - g = gevent.getcurrent() - while g is not None: - pid = id(g) - if pid in _overrides: - _root_ids[gid] = pid - return pid - g = getattr(g, "parent", None) except ImportError: - pass + # gevent is worker-only; without it there is no parent chain to walk. + return None + + g = gevent.getcurrent() + while g is not None: + pid = id(g) + if pid in _overrides: + _root_ids[gid] = pid + return pid + g = getattr(g, "parent", None) return None From 9897f04dd26e66de14da51938394c5b8829b1dbc Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 21:38:58 +0800 Subject: [PATCH 16/19] fix: use partial per-channel BYOK overrides Text-only or vision-only llm_config no longer cross-fills the other channel, so text models cannot break VLM calls (and vice versa). Co-authored-by: Cursor --- apps/api/app/api/v2/routes/retrieval.py | 7 +- .../shared/models/schemas/job.py | 7 +- .../shared/models/schemas/llm_config.py | 19 +-- scripts/smoke_byok.sh | 110 ++++++++++++++++++ 4 files changed, 129 insertions(+), 14 deletions(-) create mode 100755 scripts/smoke_byok.sh diff --git a/apps/api/app/api/v2/routes/retrieval.py b/apps/api/app/api/v2/routes/retrieval.py index 0516fb32d..bc69d7677 100644 --- a/apps/api/app/api/v2/routes/retrieval.py +++ b/apps/api/app/api/v2/routes/retrieval.py @@ -26,9 +26,10 @@ class RetrievalQueryRequestV2(RetrievalQueryRequest): None, description=( "Optional bring-your-own-key OpenAI-compatible LLM credentials. " - "Provide text and/or vision provider configs; if only one is set it " - "is used as a unified multimodal model for both. When omitted, " - "server defaults are used." + "Provide text and/or vision provider configs; each slot overrides " + "only its own channel (missing slots keep server defaults). " + "To use one multimodal model for both, set text and vision to the " + "same credentials. When omitted, server defaults are used." ), ) diff --git a/packages/shared-python/shared/models/schemas/job.py b/packages/shared-python/shared/models/schemas/job.py index c772bf680..6edd45909 100644 --- a/packages/shared-python/shared/models/schemas/job.py +++ b/packages/shared-python/shared/models/schemas/job.py @@ -91,9 +91,10 @@ class JobCreateV2(JobCreateBase): None, description=( "Optional bring-your-own-key OpenAI-compatible LLM credentials. " - "Provide text and/or vision provider configs; if only one is set it " - "is used as a unified multimodal model for both. When omitted, " - "server defaults are used." + "Provide text and/or vision provider configs; each slot overrides " + "only its own channel (missing slots keep server defaults). " + "To use one multimodal model for both, set text and vision to the " + "same credentials. When omitted, server defaults are used." ), ) diff --git a/packages/shared-python/shared/models/schemas/llm_config.py b/packages/shared-python/shared/models/schemas/llm_config.py index 4b71cfe4d..1d85f36a4 100644 --- a/packages/shared-python/shared/models/schemas/llm_config.py +++ b/packages/shared-python/shared/models/schemas/llm_config.py @@ -22,11 +22,14 @@ class LLMProviderConfig(BaseModel): class LLMConfig(BaseModel): """Optional text + vision provider configs for BYOK. - Semantics: - - both set -> text for text tasks, vision for VLM tasks - - exactly one set -> that config is used as a unified multimodal model - for both text and vision tasks + Semantics (partial override per channel): + - ``text`` set -> overrides text / planning LLM calls only + - ``vision`` set -> overrides vision / VLM calls only + - a missing slot keeps the server default for that channel - neither set -> invalid when the object itself is present + + To drive both channels with one multimodal model, set both ``text`` and + ``vision`` to the same credentials. """ text: Optional[LLMProviderConfig] = Field( @@ -43,12 +46,12 @@ def _require_at_least_one_provider(self) -> "LLMConfig": return self def text_effective(self) -> LLMProviderConfig | None: - """Resolve the config used for text tasks (unified-multimodal aware).""" - return self.text if self.text is not None else self.vision + """Return the text-channel override, or None to keep server defaults.""" + return self.text def vision_effective(self) -> LLMProviderConfig | None: - """Resolve the config used for vision/VLM tasks (unified-multimodal aware).""" - return self.vision if self.vision is not None else self.text + """Return the vision-channel override, or None to keep server defaults.""" + return self.vision def masked_dump(self) -> dict[str, Any]: """Serialize with api_key values redacted for snapshots / responses.""" diff --git a/scripts/smoke_byok.sh b/scripts/smoke_byok.sh new file mode 100755 index 000000000..0ea3be8aa --- /dev/null +++ b/scripts/smoke_byok.sh @@ -0,0 +1,110 @@ +#!/usr/bin/env bash +# Local BYOK smoke checks against a running KNOWHERE API on :5005. +# Usage: +# export KNOWHERE_API_KEY=kh_... +# ./scripts/smoke_byok.sh +set -euo pipefail + +BASE_URL="${KNOWHERE_BASE_URL:-http://localhost:5005}" +API_KEY="${KNOWHERE_API_KEY:?Set KNOWHERE_API_KEY to a local API key}" + +auth_hdr=(-H "Authorization: Bearer ${API_KEY}" -H "Content-Type: application/json") + +echo "== health ==" +curl -fsS "${BASE_URL}/health" | head -c 200; echo + +echo "== OpenAPI: v2 jobs has llm_config, v1 does not ==" +python3 - <<'PY' +import json, os, urllib.request +base = os.environ.get("KNOWHERE_BASE_URL", "http://localhost:5005") +with urllib.request.urlopen(f"{base}/openapi.json") as resp: + schema = json.load(resp) +paths = schema["paths"] +v1 = paths.get("/api/v1/jobs", paths.get("/v1/jobs", {})) +v2 = paths.get("/api/v2/jobs", paths.get("/v2/jobs", {})) +# FastAPI root_path may strip /api from openapi paths — try both +def body_props(op): + if not op: + return set() + post = op.get("post") or {} + body = ((post.get("requestBody") or {}).get("content") or {}).get("application/json") or {} + s = body.get("schema") or {} + if "$ref" in s: + name = s["$ref"].rsplit("/", 1)[-1] + s = schema["components"]["schemas"].get(name, {}) + return set((s.get("properties") or {}).keys()) + +# Prefer component schemas when available +comps = schema.get("components", {}).get("schemas", {}) +v1_keys = set((comps.get("JobCreate") or {}).get("properties", {}).keys()) +v2_keys = set((comps.get("JobCreateV2") or {}).get("properties", {}).keys()) +r1 = set((comps.get("RetrievalQueryRequest") or {}).get("properties", {}).keys()) +r2 = set((comps.get("RetrievalQueryRequestV2") or {}).get("properties", {}).keys()) +print("JobCreate llm_config:", "llm_config" in v1_keys) +print("JobCreateV2 llm_config:", "llm_config" in v2_keys) +print("RetrievalQueryRequest llm_config:", "llm_config" in r1) +print("RetrievalQueryRequestV2 llm_config:", "llm_config" in r2) +assert "llm_config" not in v1_keys +assert "llm_config" in v2_keys +assert "llm_config" not in r1 +assert "llm_config" in r2 +print("OpenAPI BYOK surface OK") +PY + +echo "== v2 jobs: reject empty llm_config object ==" +code=$(curl -sS -o /tmp/byok_empty.json -w '%{http_code}' "${auth_hdr[@]}" \ + -d '{"source_type":"url","source_url":"https://example.com/a.pdf","llm_config":{}}' \ + "${BASE_URL}/api/v2/jobs") +echo "HTTP $code" +cat /tmp/byok_empty.json | head -c 400; echo +test "$code" = "422" || test "$code" = "400" + +echo "== v2 jobs: accept llm_config.text-only (partial override) shape ==" +# Expect waiting-file or pending-ish success, not 422 +code=$(curl -sS -o /tmp/byok_ok.json -w '%{http_code}' "${auth_hdr[@]}" \ + -d '{ + "source_type":"url", + "source_url":"https://www.w3.org/WAI/ER/tests/xhtml/testfiles/resources/pdf/dummy.pdf", + "file_name":"dummy.pdf", + "data_id":"byok-smoke-text-only", + "llm_config":{ + "text":{ + "api_key":"sk-smoke-test-key-not-real", + "model":"gpt-4o-mini", + "base_url":"https://api.openai.com/v1" + } + } + }' \ + "${BASE_URL}/api/v2/jobs") +echo "HTTP $code" +cat /tmp/byok_ok.json | head -c 600; echo +test "$code" = "200" || test "$code" = "201" + +JOB_ID=$(python3 -c 'import json; print(json.load(open("/tmp/byok_ok.json")).get("job_id",""))') +echo "job_id=${JOB_ID}" + +if [[ -n "${JOB_ID}" ]]; then + echo "== GET job: ensure raw api_key is not echoed ==" + curl -fsS "${auth_hdr[@]}" "${BASE_URL}/api/v2/jobs/${JOB_ID}" | tee /tmp/byok_job.json | head -c 800; echo + if grep -q 'sk-smoke-test-key-not-real' /tmp/byok_job.json; then + echo "FAIL: raw API key leaked in job response" >&2 + exit 1 + fi + echo "No raw key in job response OK" +fi + +echo "== v1 jobs: llm_config should be rejected as unsupported extra ==" +code=$(curl -sS -o /tmp/byok_v1.json -w '%{http_code}' "${auth_hdr[@]}" \ + -d '{ + "source_type":"url", + "source_url":"https://example.com/a.pdf", + "llm_config":{"text":{"api_key":"sk-x","model":"m","base_url":"https://example.com/v1"}} + }' \ + "${BASE_URL}/api/v1/jobs") +echo "HTTP $code" +cat /tmp/byok_v1.json | head -c 400; echo +# Prefer 422 / validation error; 200 would mean v1 accepted BYOK (bad) +test "$code" != "200" && test "$code" != "201" + +echo +echo "BYOK smoke checks finished." From f214091e4ababe260622cb09ae582537d3fb4eea Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 21:44:59 +0800 Subject: [PATCH 17/19] feat: add llm_config.provider multimodal shorthand One shared provider slot covers both text and vision; text/vision still override per channel when set. Co-authored-by: Cursor --- apps/api/app/api/v2/routes/retrieval.py | 8 +- .../shared/models/schemas/job.py | 8 +- .../shared/models/schemas/job_metadata.py | 2 +- .../shared/models/schemas/llm_config.py | 40 ++++++---- .../shared/tests/test_llm_config.py | 75 +++++++++++++++++++ scripts/smoke_byok.sh | 28 ++++++- 6 files changed, 134 insertions(+), 27 deletions(-) create mode 100644 packages/shared-python/shared/tests/test_llm_config.py diff --git a/apps/api/app/api/v2/routes/retrieval.py b/apps/api/app/api/v2/routes/retrieval.py index bc69d7677..0aca5e36d 100644 --- a/apps/api/app/api/v2/routes/retrieval.py +++ b/apps/api/app/api/v2/routes/retrieval.py @@ -26,10 +26,10 @@ class RetrievalQueryRequestV2(RetrievalQueryRequest): None, description=( "Optional bring-your-own-key OpenAI-compatible LLM credentials. " - "Provide text and/or vision provider configs; each slot overrides " - "only its own channel (missing slots keep server defaults). " - "To use one multimodal model for both, set text and vision to the " - "same credentials. When omitted, server defaults are used." + "Use `provider` for a single multimodal model (both channels). " + "Use `text` / `vision` to override one channel (or both). " + "Missing channels without `provider` keep server defaults. " + "When omitted, server defaults are used." ), ) diff --git a/packages/shared-python/shared/models/schemas/job.py b/packages/shared-python/shared/models/schemas/job.py index 6edd45909..adf418758 100644 --- a/packages/shared-python/shared/models/schemas/job.py +++ b/packages/shared-python/shared/models/schemas/job.py @@ -91,10 +91,10 @@ class JobCreateV2(JobCreateBase): None, description=( "Optional bring-your-own-key OpenAI-compatible LLM credentials. " - "Provide text and/or vision provider configs; each slot overrides " - "only its own channel (missing slots keep server defaults). " - "To use one multimodal model for both, set text and vision to the " - "same credentials. When omitted, server defaults are used." + "Use `provider` for a single multimodal model (both channels). " + "Use `text` / `vision` to override one channel (or both). " + "Missing channels without `provider` keep server defaults. " + "When omitted, server defaults are used." ), ) diff --git a/packages/shared-python/shared/models/schemas/job_metadata.py b/packages/shared-python/shared/models/schemas/job_metadata.py index 9e25a2d7a..026ed8f1d 100644 --- a/packages/shared-python/shared/models/schemas/job_metadata.py +++ b/packages/shared-python/shared/models/schemas/job_metadata.py @@ -268,7 +268,7 @@ def _mask_llm_config_in_request(payload: Dict[str, Any]) -> Dict[str, Any]: masked = dict(payload) masked_llm: Dict[str, Any] = {} - for slot in ("text", "vision"): + for slot in ("provider", "text", "vision"): provider = llm_config.get(slot) if isinstance(provider, dict): provider_copy = dict(provider) diff --git a/packages/shared-python/shared/models/schemas/llm_config.py b/packages/shared-python/shared/models/schemas/llm_config.py index 1d85f36a4..2e71362c9 100644 --- a/packages/shared-python/shared/models/schemas/llm_config.py +++ b/packages/shared-python/shared/models/schemas/llm_config.py @@ -22,36 +22,47 @@ class LLMProviderConfig(BaseModel): class LLMConfig(BaseModel): """Optional text + vision provider configs for BYOK. - Semantics (partial override per channel): - - ``text`` set -> overrides text / planning LLM calls only - - ``vision`` set -> overrides vision / VLM calls only - - a missing slot keeps the server default for that channel - - neither set -> invalid when the object itself is present - - To drive both channels with one multimodal model, set both ``text`` and - ``vision`` to the same credentials. + Semantics: + - ``provider`` set -> baseline for both text and vision channels + - ``text`` / ``vision`` override that channel (and win over ``provider``) + - a channel with neither a slot nor ``provider`` keeps server defaults + - none of ``provider`` / ``text`` / ``vision`` set -> invalid when present + + Multimodal shorthand (one model for both channels):: + + {"provider": {"api_key": "...", "model": "gpt-4o", "base_url": "..."}} """ + provider: Optional[LLMProviderConfig] = Field( + None, + description=( + "Shared OpenAI-compatible credentials for both text and vision. " + "Use this for a single multimodal model; override with text/vision " + "when channels need different endpoints." + ), + ) text: Optional[LLMProviderConfig] = Field( - None, description="Text / planning LLM credentials" + None, description="Text / planning LLM credentials (overrides provider)" ) vision: Optional[LLMProviderConfig] = Field( - None, description="Vision / VLM credentials" + None, description="Vision / VLM credentials (overrides provider)" ) @model_validator(mode="after") def _require_at_least_one_provider(self) -> "LLMConfig": - if self.text is None and self.vision is None: - raise ValueError("llm_config requires at least one of text or vision") + if self.provider is None and self.text is None and self.vision is None: + raise ValueError( + "llm_config requires at least one of provider, text, or vision" + ) return self def text_effective(self) -> LLMProviderConfig | None: """Return the text-channel override, or None to keep server defaults.""" - return self.text + return self.text if self.text is not None else self.provider def vision_effective(self) -> LLMProviderConfig | None: """Return the vision-channel override, or None to keep server defaults.""" - return self.vision + return self.vision if self.vision is not None else self.provider def masked_dump(self) -> dict[str, Any]: """Serialize with api_key values redacted for snapshots / responses.""" @@ -67,6 +78,7 @@ def _mask_provider(provider: LLMProviderConfig | None) -> dict[str, Any] | None: } return { + "provider": _mask_provider(self.provider), "text": _mask_provider(self.text), "vision": _mask_provider(self.vision), } diff --git a/packages/shared-python/shared/tests/test_llm_config.py b/packages/shared-python/shared/tests/test_llm_config.py new file mode 100644 index 000000000..0ea54bd6c --- /dev/null +++ b/packages/shared-python/shared/tests/test_llm_config.py @@ -0,0 +1,75 @@ +"""Unit tests for BYOK LLMConfig resolution.""" + +from __future__ import annotations + +import os + +import pytest +from pydantic import ValidationError + +os.environ.setdefault("DATABASE_URL", "postgresql+asyncpg://test:test@localhost/test") +os.environ.setdefault("TMP_PATH", "/tmp/knowhere-test") +os.environ.setdefault("S3_BUCKET_NAME", "test-uploads") +os.environ.setdefault("S3_ACCESS_KEY_ID", "test") +os.environ.setdefault("S3_SECRET_ACCESS_KEY", "test") +os.environ.setdefault("S3_TEMP_PATH", "/tmp") + +from shared.models.schemas.llm_config import LLMConfig, LLMProviderConfig + + +def _creds(model: str = "gpt-4o") -> LLMProviderConfig: + return LLMProviderConfig( + api_key="sk-test", + model=model, + base_url="https://api.openai.com/v1", + ) + + +def test_provider_alone_applies_to_both_channels() -> None: + cfg = LLMConfig(provider=_creds("gpt-4o")) + assert cfg.text_effective() is not None + assert cfg.vision_effective() is not None + assert cfg.text_effective().model == "gpt-4o" + assert cfg.vision_effective().model == "gpt-4o" + + +def test_text_only_leaves_vision_on_defaults() -> None: + cfg = LLMConfig(text=_creds("text-model")) + assert cfg.text_effective().model == "text-model" + assert cfg.vision_effective() is None + + +def test_vision_only_leaves_text_on_defaults() -> None: + cfg = LLMConfig(vision=_creds("vlm")) + assert cfg.text_effective() is None + assert cfg.vision_effective().model == "vlm" + + +def test_channel_overrides_provider() -> None: + cfg = LLMConfig( + provider=_creds("shared"), + text=_creds("text-only"), + vision=_creds("vision-only"), + ) + assert cfg.text_effective().model == "text-only" + assert cfg.vision_effective().model == "vision-only" + + +def test_provider_plus_text_override() -> None: + cfg = LLMConfig(provider=_creds("shared"), text=_creds("text-only")) + assert cfg.text_effective().model == "text-only" + assert cfg.vision_effective().model == "shared" + + +def test_empty_config_rejected() -> None: + with pytest.raises(ValidationError, match="provider, text, or vision"): + LLMConfig() + + +def test_masked_dump_includes_provider() -> None: + cfg = LLMConfig(provider=_creds()) + dump = cfg.masked_dump() + assert dump["provider"] is not None + assert dump["provider"]["api_key"] != "sk-test" + assert dump["text"] is None + assert dump["vision"] is None diff --git a/scripts/smoke_byok.sh b/scripts/smoke_byok.sh index 0ea3be8aa..0bc0bac01 100755 --- a/scripts/smoke_byok.sh +++ b/scripts/smoke_byok.sh @@ -59,18 +59,18 @@ echo "HTTP $code" cat /tmp/byok_empty.json | head -c 400; echo test "$code" = "422" || test "$code" = "400" -echo "== v2 jobs: accept llm_config.text-only (partial override) shape ==" +echo "== v2 jobs: accept llm_config.provider (multimodal shorthand) shape ==" # Expect waiting-file or pending-ish success, not 422 code=$(curl -sS -o /tmp/byok_ok.json -w '%{http_code}' "${auth_hdr[@]}" \ -d '{ "source_type":"url", "source_url":"https://www.w3.org/WAI/ER/tests/xhtml/testfiles/resources/pdf/dummy.pdf", "file_name":"dummy.pdf", - "data_id":"byok-smoke-text-only", + "data_id":"byok-smoke-provider", "llm_config":{ - "text":{ + "provider":{ "api_key":"sk-smoke-test-key-not-real", - "model":"gpt-4o-mini", + "model":"gpt-4o", "base_url":"https://api.openai.com/v1" } } @@ -93,6 +93,26 @@ if [[ -n "${JOB_ID}" ]]; then echo "No raw key in job response OK" fi +echo "== v2 jobs: accept llm_config.text-only (partial override) shape ==" +code=$(curl -sS -o /tmp/byok_text.json -w '%{http_code}' "${auth_hdr[@]}" \ + -d '{ + "source_type":"url", + "source_url":"https://www.w3.org/WAI/ER/tests/xhtml/testfiles/resources/pdf/dummy.pdf", + "file_name":"dummy.pdf", + "data_id":"byok-smoke-text-only", + "llm_config":{ + "text":{ + "api_key":"sk-smoke-test-key-not-real", + "model":"gpt-4o-mini", + "base_url":"https://api.openai.com/v1" + } + } + }' \ + "${BASE_URL}/api/v2/jobs") +echo "HTTP $code" +cat /tmp/byok_text.json | head -c 400; echo +test "$code" = "200" || test "$code" = "201" + echo "== v1 jobs: llm_config should be rejected as unsupported extra ==" code=$(curl -sS -o /tmp/byok_v1.json -w '%{http_code}' "${auth_hdr[@]}" \ -d '{ From b72803e4e10fce9d39d6f93ab998cadb4ec9b77b Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 22:03:46 +0800 Subject: [PATCH 18/19] refactor: flatten llm_config to OpenAI-style root fields Use flat api_key/model/base_url for multimodal BYOK; keep text/vision as full channel replacements for different provider endpoints. Co-authored-by: Cursor --- apps/api/app/api/v2/routes/retrieval.py | 6 +- .../shared/models/schemas/job.py | 6 +- .../shared/models/schemas/job_metadata.py | 10 +-- .../shared/models/schemas/llm_config.py | 78 +++++++++++++------ .../shared/tests/test_llm_config.py | 77 +++++++++++++----- scripts/smoke_byok.sh | 37 +++++++-- 6 files changed, 151 insertions(+), 63 deletions(-) diff --git a/apps/api/app/api/v2/routes/retrieval.py b/apps/api/app/api/v2/routes/retrieval.py index 0aca5e36d..2b04a83c2 100644 --- a/apps/api/app/api/v2/routes/retrieval.py +++ b/apps/api/app/api/v2/routes/retrieval.py @@ -26,9 +26,9 @@ class RetrievalQueryRequestV2(RetrievalQueryRequest): None, description=( "Optional bring-your-own-key OpenAI-compatible LLM credentials. " - "Use `provider` for a single multimodal model (both channels). " - "Use `text` / `vision` to override one channel (or both). " - "Missing channels without `provider` keep server defaults. " + "Flat root {api_key, model, base_url} applies to both channels " + "(multimodal shorthand). Optional text/vision objects fully replace " + "the default for that channel (use both for different endpoints). " "When omitted, server defaults are used." ), ) diff --git a/packages/shared-python/shared/models/schemas/job.py b/packages/shared-python/shared/models/schemas/job.py index adf418758..effa050e8 100644 --- a/packages/shared-python/shared/models/schemas/job.py +++ b/packages/shared-python/shared/models/schemas/job.py @@ -91,9 +91,9 @@ class JobCreateV2(JobCreateBase): None, description=( "Optional bring-your-own-key OpenAI-compatible LLM credentials. " - "Use `provider` for a single multimodal model (both channels). " - "Use `text` / `vision` to override one channel (or both). " - "Missing channels without `provider` keep server defaults. " + "Flat root {api_key, model, base_url} applies to both channels " + "(multimodal shorthand). Optional text/vision objects fully replace " + "the default for that channel (use both for different endpoints). " "When omitted, server defaults are used." ), ) diff --git a/packages/shared-python/shared/models/schemas/job_metadata.py b/packages/shared-python/shared/models/schemas/job_metadata.py index 026ed8f1d..4efb808b3 100644 --- a/packages/shared-python/shared/models/schemas/job_metadata.py +++ b/packages/shared-python/shared/models/schemas/job_metadata.py @@ -267,9 +267,11 @@ def _mask_llm_config_in_request(payload: Dict[str, Any]) -> Dict[str, Any]: return payload masked = dict(payload) - masked_llm: Dict[str, Any] = {} - for slot in ("provider", "text", "vision"): - provider = llm_config.get(slot) + masked_llm: Dict[str, Any] = dict(llm_config) + if isinstance(masked_llm.get("api_key"), str): + masked_llm["api_key"] = mask_api_key(masked_llm["api_key"]) + for slot in ("text", "vision"): + provider = masked_llm.get(slot) if isinstance(provider, dict): provider_copy = dict(provider) if "api_key" in provider_copy: @@ -279,8 +281,6 @@ def _mask_llm_config_in_request(payload: Dict[str, Any]) -> Dict[str, Any]: else None ) masked_llm[slot] = provider_copy - elif provider is not None: - masked_llm[slot] = provider masked["llm_config"] = masked_llm return masked diff --git a/packages/shared-python/shared/models/schemas/llm_config.py b/packages/shared-python/shared/models/schemas/llm_config.py index 2e71362c9..fd7e7f2bb 100644 --- a/packages/shared-python/shared/models/schemas/llm_config.py +++ b/packages/shared-python/shared/models/schemas/llm_config.py @@ -20,49 +20,74 @@ class LLMProviderConfig(BaseModel): class LLMConfig(BaseModel): - """Optional text + vision provider configs for BYOK. + """OpenAI-compatible BYOK credentials (flat root + optional channel overrides). - Semantics: - - ``provider`` set -> baseline for both text and vision channels - - ``text`` / ``vision`` override that channel (and win over ``provider``) - - a channel with neither a slot nor ``provider`` keeps server defaults - - none of ``provider`` / ``text`` / ``vision`` set -> invalid when present + Happy path (one multimodal model for both channels), matching OpenAI / + LangChain / LiteLLM style:: + + {"api_key": "...", "model": "gpt-4o", "base_url": "https://api.openai.com/v1"} + + Different endpoints per channel:: - Multimodal shorthand (one model for both channels):: + { + "text": {"api_key": "...", "model": "...", "base_url": "..."}, + "vision": {"api_key": "...", "model": "...", "base_url": "..."} + } - {"provider": {"api_key": "...", "model": "gpt-4o", "base_url": "..."}} + Semantics: + - root ``api_key`` / ``model`` / ``base_url`` (all three together) -> default + for both channels + - ``text`` / ``vision`` fully replace the default for that channel + - a channel with neither a slot nor a root default keeps server defaults """ - provider: Optional[LLMProviderConfig] = Field( + api_key: Optional[str] = Field(None, min_length=1, description="Default provider API key") + model: Optional[str] = Field(None, min_length=1, description="Default model identifier") + base_url: Optional[str] = Field( None, - description=( - "Shared OpenAI-compatible credentials for both text and vision. " - "Use this for a single multimodal model; override with text/vision " - "when channels need different endpoints." - ), + min_length=1, + description="Default OpenAI-compatible base URL", ) text: Optional[LLMProviderConfig] = Field( - None, description="Text / planning LLM credentials (overrides provider)" + None, + description="Text / planning credentials (replaces root for text channel)", ) vision: Optional[LLMProviderConfig] = Field( - None, description="Vision / VLM credentials (overrides provider)" + None, + description="Vision / VLM credentials (replaces root for vision channel)", ) @model_validator(mode="after") - def _require_at_least_one_provider(self) -> "LLMConfig": - if self.provider is None and self.text is None and self.vision is None: + def _validate_shape(self) -> "LLMConfig": + root_fields = (self.api_key, self.model, self.base_url) + root_set_count = sum(value is not None for value in root_fields) + if root_set_count not in (0, 3): raise ValueError( - "llm_config requires at least one of provider, text, or vision" + "llm_config root api_key, model, and base_url must be set together" + ) + if root_set_count == 0 and self.text is None and self.vision is None: + raise ValueError( + "llm_config requires root credentials and/or text/vision overrides" ) return self + def root_provider(self) -> LLMProviderConfig | None: + """Return the flat root as a provider config, or None if unset.""" + if self.api_key is None or self.model is None or self.base_url is None: + return None + return LLMProviderConfig( + api_key=self.api_key, + model=self.model, + base_url=self.base_url, + ) + def text_effective(self) -> LLMProviderConfig | None: - """Return the text-channel override, or None to keep server defaults.""" - return self.text if self.text is not None else self.provider + """Return the text-channel config, or None to keep server defaults.""" + return self.text if self.text is not None else self.root_provider() def vision_effective(self) -> LLMProviderConfig | None: - """Return the vision-channel override, or None to keep server defaults.""" - return self.vision if self.vision is not None else self.provider + """Return the vision-channel config, or None to keep server defaults.""" + return self.vision if self.vision is not None else self.root_provider() def masked_dump(self) -> dict[str, Any]: """Serialize with api_key values redacted for snapshots / responses.""" @@ -77,11 +102,14 @@ def _mask_provider(provider: LLMProviderConfig | None) -> dict[str, Any] | None: "base_url": provider.base_url, } - return { - "provider": _mask_provider(self.provider), + dump: dict[str, Any] = { + "api_key": mask_api_key(self.api_key) if self.api_key else None, + "model": self.model, + "base_url": self.base_url, "text": _mask_provider(self.text), "vision": _mask_provider(self.vision), } + return dump def parse_llm_config(value: Any) -> LLMConfig | None: diff --git a/packages/shared-python/shared/tests/test_llm_config.py b/packages/shared-python/shared/tests/test_llm_config.py index 0ea54bd6c..5ad0b377c 100644 --- a/packages/shared-python/shared/tests/test_llm_config.py +++ b/packages/shared-python/shared/tests/test_llm_config.py @@ -17,20 +17,25 @@ from shared.models.schemas.llm_config import LLMConfig, LLMProviderConfig -def _creds(model: str = "gpt-4o") -> LLMProviderConfig: - return LLMProviderConfig( - api_key="sk-test", - model=model, - base_url="https://api.openai.com/v1", - ) +def _creds( + model: str = "gpt-4o", + *, + api_key: str = "sk-test", + base_url: str = "https://api.openai.com/v1", +) -> LLMProviderConfig: + return LLMProviderConfig(api_key=api_key, model=model, base_url=base_url) -def test_provider_alone_applies_to_both_channels() -> None: - cfg = LLMConfig(provider=_creds("gpt-4o")) +def test_flat_root_applies_to_both_channels() -> None: + cfg = LLMConfig( + api_key="sk-root", + model="gpt-4o", + base_url="https://api.openai.com/v1", + ) assert cfg.text_effective() is not None assert cfg.vision_effective() is not None assert cfg.text_effective().model == "gpt-4o" - assert cfg.vision_effective().model == "gpt-4o" + assert cfg.vision_effective().api_key == "sk-root" def test_text_only_leaves_vision_on_defaults() -> None: @@ -45,31 +50,63 @@ def test_vision_only_leaves_text_on_defaults() -> None: assert cfg.vision_effective().model == "vlm" -def test_channel_overrides_provider() -> None: +def test_two_different_endpoints() -> None: cfg = LLMConfig( - provider=_creds("shared"), + text=_creds("gpt-4o-mini", base_url="https://api.openai.com/v1"), + vision=_creds( + "qwen-vl-max", + api_key="sk-ali", + base_url="https://dashscope.aliyuncs.com/compatible-mode/v1", + ), + ) + assert cfg.text_effective().base_url == "https://api.openai.com/v1" + assert cfg.vision_effective().base_url.endswith("/compatible-mode/v1") + assert cfg.vision_effective().model == "qwen-vl-max" + + +def test_channel_replaces_root() -> None: + cfg = LLMConfig( + api_key="sk-root", + model="shared", + base_url="https://api.openai.com/v1", text=_creds("text-only"), - vision=_creds("vision-only"), + vision=_creds("vision-only", base_url="https://other.example/v1"), ) assert cfg.text_effective().model == "text-only" assert cfg.vision_effective().model == "vision-only" + assert cfg.vision_effective().base_url == "https://other.example/v1" -def test_provider_plus_text_override() -> None: - cfg = LLMConfig(provider=_creds("shared"), text=_creds("text-only")) +def test_root_plus_text_override() -> None: + cfg = LLMConfig( + api_key="sk-root", + model="shared", + base_url="https://api.openai.com/v1", + text=_creds("text-only"), + ) assert cfg.text_effective().model == "text-only" assert cfg.vision_effective().model == "shared" +def test_partial_root_rejected() -> None: + with pytest.raises(ValidationError, match="must be set together"): + LLMConfig(api_key="sk-only") + + def test_empty_config_rejected() -> None: - with pytest.raises(ValidationError, match="provider, text, or vision"): + with pytest.raises(ValidationError, match="root credentials and/or text/vision"): LLMConfig() -def test_masked_dump_includes_provider() -> None: - cfg = LLMConfig(provider=_creds()) +def test_masked_dump_masks_root_and_channels() -> None: + cfg = LLMConfig( + api_key="sk-root", + model="gpt-4o", + base_url="https://api.openai.com/v1", + vision=_creds("vlm", api_key="sk-vision"), + ) dump = cfg.masked_dump() - assert dump["provider"] is not None - assert dump["provider"]["api_key"] != "sk-test" + assert dump["api_key"] != "sk-root" + assert dump["model"] == "gpt-4o" + assert dump["vision"]["api_key"] != "sk-vision" assert dump["text"] is None - assert dump["vision"] is None diff --git a/scripts/smoke_byok.sh b/scripts/smoke_byok.sh index 0bc0bac01..104848f78 100755 --- a/scripts/smoke_byok.sh +++ b/scripts/smoke_byok.sh @@ -59,20 +59,18 @@ echo "HTTP $code" cat /tmp/byok_empty.json | head -c 400; echo test "$code" = "422" || test "$code" = "400" -echo "== v2 jobs: accept llm_config.provider (multimodal shorthand) shape ==" +echo "== v2 jobs: accept flat llm_config (multimodal shorthand) shape ==" # Expect waiting-file or pending-ish success, not 422 code=$(curl -sS -o /tmp/byok_ok.json -w '%{http_code}' "${auth_hdr[@]}" \ -d '{ "source_type":"url", "source_url":"https://www.w3.org/WAI/ER/tests/xhtml/testfiles/resources/pdf/dummy.pdf", "file_name":"dummy.pdf", - "data_id":"byok-smoke-provider", + "data_id":"byok-smoke-flat", "llm_config":{ - "provider":{ - "api_key":"sk-smoke-test-key-not-real", - "model":"gpt-4o", - "base_url":"https://api.openai.com/v1" - } + "api_key":"sk-smoke-test-key-not-real", + "model":"gpt-4o", + "base_url":"https://api.openai.com/v1" } }' \ "${BASE_URL}/api/v2/jobs") @@ -93,6 +91,31 @@ if [[ -n "${JOB_ID}" ]]; then echo "No raw key in job response OK" fi +echo "== v2 jobs: accept llm_config.text + vision (two endpoints) shape ==" +code=$(curl -sS -o /tmp/byok_split.json -w '%{http_code}' "${auth_hdr[@]}" \ + -d '{ + "source_type":"url", + "source_url":"https://www.w3.org/WAI/ER/tests/xhtml/testfiles/resources/pdf/dummy.pdf", + "file_name":"dummy.pdf", + "data_id":"byok-smoke-split", + "llm_config":{ + "text":{ + "api_key":"sk-smoke-text", + "model":"gpt-4o-mini", + "base_url":"https://api.openai.com/v1" + }, + "vision":{ + "api_key":"sk-smoke-vision", + "model":"qwen-vl-max", + "base_url":"https://dashscope.aliyuncs.com/compatible-mode/v1" + } + } + }' \ + "${BASE_URL}/api/v2/jobs") +echo "HTTP $code" +cat /tmp/byok_split.json | head -c 400; echo +test "$code" = "200" || test "$code" = "201" + echo "== v2 jobs: accept llm_config.text-only (partial override) shape ==" code=$(curl -sS -o /tmp/byok_text.json -w '%{http_code}' "${auth_hdr[@]}" \ -d '{ From c04925e62e138024470748cb9027410a8b063b35 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 15 Jul 2026 22:11:01 +0800 Subject: [PATCH 19/19] feat: support llm_config.models per-channel model map Allow shared api_key/base_url with models.text / models.vision when only the model id differs between channels. Co-authored-by: Cursor --- apps/api/app/api/v2/routes/retrieval.py | 6 +- .../shared/models/schemas/job.py | 6 +- .../shared/models/schemas/llm_config.py | 91 +++++++++++++++---- .../shared/tests/test_llm_config.py | 42 ++++++++- 4 files changed, 117 insertions(+), 28 deletions(-) diff --git a/apps/api/app/api/v2/routes/retrieval.py b/apps/api/app/api/v2/routes/retrieval.py index 2b04a83c2..6870c72f2 100644 --- a/apps/api/app/api/v2/routes/retrieval.py +++ b/apps/api/app/api/v2/routes/retrieval.py @@ -26,9 +26,9 @@ class RetrievalQueryRequestV2(RetrievalQueryRequest): None, description=( "Optional bring-your-own-key OpenAI-compatible LLM credentials. " - "Flat root {api_key, model, base_url} applies to both channels " - "(multimodal shorthand). Optional text/vision objects fully replace " - "the default for that channel (use both for different endpoints). " + "Flat {api_key, model, base_url} applies to both channels. " + "Use models.{text,vision} for different model ids on the same " + "endpoint, or text/vision objects for different endpoints. " "When omitted, server defaults are used." ), ) diff --git a/packages/shared-python/shared/models/schemas/job.py b/packages/shared-python/shared/models/schemas/job.py index effa050e8..2a948f792 100644 --- a/packages/shared-python/shared/models/schemas/job.py +++ b/packages/shared-python/shared/models/schemas/job.py @@ -91,9 +91,9 @@ class JobCreateV2(JobCreateBase): None, description=( "Optional bring-your-own-key OpenAI-compatible LLM credentials. " - "Flat root {api_key, model, base_url} applies to both channels " - "(multimodal shorthand). Optional text/vision objects fully replace " - "the default for that channel (use both for different endpoints). " + "Flat {api_key, model, base_url} applies to both channels. " + "Use models.{text,vision} for different model ids on the same " + "endpoint, or text/vision objects for different endpoints. " "When omitted, server defaults are used." ), ) diff --git a/packages/shared-python/shared/models/schemas/llm_config.py b/packages/shared-python/shared/models/schemas/llm_config.py index fd7e7f2bb..cdcc3d4bd 100644 --- a/packages/shared-python/shared/models/schemas/llm_config.py +++ b/packages/shared-python/shared/models/schemas/llm_config.py @@ -19,14 +19,28 @@ class LLMProviderConfig(BaseModel): ) +class LLMModelsConfig(BaseModel): + """Per-channel model ids that share root api_key / base_url.""" + + text: Optional[str] = Field(None, min_length=1, description="Text / planning model id") + vision: Optional[str] = Field(None, min_length=1, description="Vision / VLM model id") + + class LLMConfig(BaseModel): """OpenAI-compatible BYOK credentials (flat root + optional channel overrides). - Happy path (one multimodal model for both channels), matching OpenAI / - LangChain / LiteLLM style:: + Happy path (one multimodal model for both channels):: {"api_key": "...", "model": "gpt-4o", "base_url": "https://api.openai.com/v1"} + Same endpoint, different models per channel:: + + { + "api_key": "...", + "base_url": "https://api.openai.com/v1", + "models": {"text": "gpt-4o-mini", "vision": "gpt-4o"} + } + Different endpoints per channel:: { @@ -35,19 +49,27 @@ class LLMConfig(BaseModel): } Semantics: - - root ``api_key`` / ``model`` / ``base_url`` (all three together) -> default - for both channels - - ``text`` / ``vision`` fully replace the default for that channel - - a channel with neither a slot nor a root default keeps server defaults + - root ``api_key`` + ``base_url`` with ``model`` and/or ``models`` -> shared auth + - ``models.`` wins over root ``model`` for that channel + - ``text`` / ``vision`` objects fully replace the root for that channel + - a channel with no resolved model keeps server defaults """ api_key: Optional[str] = Field(None, min_length=1, description="Default provider API key") - model: Optional[str] = Field(None, min_length=1, description="Default model identifier") + model: Optional[str] = Field( + None, + min_length=1, + description="Default model for both channels (overridden by models.*)", + ) base_url: Optional[str] = Field( None, min_length=1, description="Default OpenAI-compatible base URL", ) + models: Optional[LLMModelsConfig] = Field( + None, + description="Per-channel model ids sharing root api_key / base_url", + ) text: Optional[LLMProviderConfig] = Field( None, description="Text / planning credentials (replaces root for text channel)", @@ -59,35 +81,64 @@ class LLMConfig(BaseModel): @model_validator(mode="after") def _validate_shape(self) -> "LLMConfig": - root_fields = (self.api_key, self.model, self.base_url) - root_set_count = sum(value is not None for value in root_fields) - if root_set_count not in (0, 3): + has_api_key = self.api_key is not None + has_base_url = self.base_url is not None + if has_api_key != has_base_url: + raise ValueError("llm_config api_key and base_url must be set together") + + has_auth = has_api_key and has_base_url + has_models = self.models is not None and ( + self.models.text is not None or self.models.vision is not None + ) + if self.models is not None and not has_models: + raise ValueError("llm_config.models requires at least one of text or vision") + if self.model is not None and not has_auth: + raise ValueError("llm_config.model requires api_key and base_url") + if has_models and not has_auth: + raise ValueError("llm_config.models requires api_key and base_url") + if has_auth and self.model is None and not has_models: raise ValueError( - "llm_config root api_key, model, and base_url must be set together" + "llm_config with api_key/base_url requires model and/or models" ) - if root_set_count == 0 and self.text is None and self.vision is None: + + if ( + not has_auth + and self.text is None + and self.vision is None + ): raise ValueError( "llm_config requires root credentials and/or text/vision overrides" ) return self - def root_provider(self) -> LLMProviderConfig | None: - """Return the flat root as a provider config, or None if unset.""" - if self.api_key is None or self.model is None or self.base_url is None: + def _channel_model(self, channel: str) -> str | None: + if self.models is not None: + named = getattr(self.models, channel) + if isinstance(named, str) and named: + return named + return self.model + + def _root_provider_for(self, channel: str) -> LLMProviderConfig | None: + if self.api_key is None or self.base_url is None: + return None + model = self._channel_model(channel) + if model is None: return None return LLMProviderConfig( api_key=self.api_key, - model=self.model, + model=model, base_url=self.base_url, ) def text_effective(self) -> LLMProviderConfig | None: """Return the text-channel config, or None to keep server defaults.""" - return self.text if self.text is not None else self.root_provider() + return self.text if self.text is not None else self._root_provider_for("text") def vision_effective(self) -> LLMProviderConfig | None: """Return the vision-channel config, or None to keep server defaults.""" - return self.vision if self.vision is not None else self.root_provider() + return ( + self.vision if self.vision is not None else self._root_provider_for("vision") + ) def masked_dump(self) -> dict[str, Any]: """Serialize with api_key values redacted for snapshots / responses.""" @@ -102,14 +153,14 @@ def _mask_provider(provider: LLMProviderConfig | None) -> dict[str, Any] | None: "base_url": provider.base_url, } - dump: dict[str, Any] = { + return { "api_key": mask_api_key(self.api_key) if self.api_key else None, "model": self.model, "base_url": self.base_url, + "models": self.models.model_dump() if self.models is not None else None, "text": _mask_provider(self.text), "vision": _mask_provider(self.vision), } - return dump def parse_llm_config(value: Any) -> LLMConfig | None: diff --git a/packages/shared-python/shared/tests/test_llm_config.py b/packages/shared-python/shared/tests/test_llm_config.py index 5ad0b377c..70d7eb6d6 100644 --- a/packages/shared-python/shared/tests/test_llm_config.py +++ b/packages/shared-python/shared/tests/test_llm_config.py @@ -14,7 +14,7 @@ os.environ.setdefault("S3_SECRET_ACCESS_KEY", "test") os.environ.setdefault("S3_TEMP_PATH", "/tmp") -from shared.models.schemas.llm_config import LLMConfig, LLMProviderConfig +from shared.models.schemas.llm_config import LLMConfig, LLMModelsConfig, LLMProviderConfig def _creds( @@ -38,6 +38,39 @@ def test_flat_root_applies_to_both_channels() -> None: assert cfg.vision_effective().api_key == "sk-root" +def test_models_map_same_endpoint_different_models() -> None: + cfg = LLMConfig( + api_key="sk-root", + base_url="https://api.openai.com/v1", + models=LLMModelsConfig(text="gpt-4o-mini", vision="gpt-4o"), + ) + assert cfg.text_effective().model == "gpt-4o-mini" + assert cfg.vision_effective().model == "gpt-4o" + assert cfg.text_effective().api_key == "sk-root" + assert cfg.vision_effective().base_url == "https://api.openai.com/v1" + + +def test_models_map_partial_leaves_other_channel_default() -> None: + cfg = LLMConfig( + api_key="sk-root", + base_url="https://api.openai.com/v1", + models=LLMModelsConfig(text="gpt-4o-mini"), + ) + assert cfg.text_effective().model == "gpt-4o-mini" + assert cfg.vision_effective() is None + + +def test_models_overrides_root_model_per_channel() -> None: + cfg = LLMConfig( + api_key="sk-root", + model="gpt-4o-mini", + base_url="https://api.openai.com/v1", + models=LLMModelsConfig(vision="gpt-4o"), + ) + assert cfg.text_effective().model == "gpt-4o-mini" + assert cfg.vision_effective().model == "gpt-4o" + + def test_text_only_leaves_vision_on_defaults() -> None: cfg = LLMConfig(text=_creds("text-model")) assert cfg.text_effective().model == "text-model" @@ -89,10 +122,15 @@ def test_root_plus_text_override() -> None: def test_partial_root_rejected() -> None: - with pytest.raises(ValidationError, match="must be set together"): + with pytest.raises(ValidationError, match="api_key and base_url must be set together"): LLMConfig(api_key="sk-only") +def test_auth_without_model_rejected() -> None: + with pytest.raises(ValidationError, match="requires model and/or models"): + LLMConfig(api_key="sk-root", base_url="https://api.openai.com/v1") + + def test_empty_config_rejected() -> None: with pytest.raises(ValidationError, match="root credentials and/or text/vision"): LLMConfig()