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/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..6870c72f2 100644 --- a/apps/api/app/api/v2/routes/retrieval.py +++ b/apps/api/app/api/v2/routes/retrieval.py @@ -1,5 +1,48 @@ """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. " + "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." + ), + ) + + +@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/api/app/services/auth/dashboard_jwt_authentication_service.py b/apps/api/app/services/auth/dashboard_jwt_authentication_service.py index 297c68c71..0989460ff 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,43 @@ 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, NoReturn, cast import jwt -from jwt import PyJWKClient +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 +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, +} 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) @@ -27,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.""" @@ -40,47 +84,78 @@ 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) - user_id = payload.get("id") - if not isinstance(user_id, str) or not user_id: - raise AuthException(user_message="Token missing 'id' claim") - - 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}, + 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, ) - return payload + payload = _verify_payload_or_reject( + token=token, + key=key, + telemetry_context=telemetry_context, + ) + return _build_identity_or_reject(payload, telemetry_context) - def _get_verification_key(self, token: str) -> Any: - """Resolve the JWT verification key from the Dashboard JWKS endpoint.""" + 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: - 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}" - ) + key = self._get_verification_key(key_id) + except PyJWKClientConnectionError as error: + _reject_jwks_dependency( + failure_reason="jwks_unavailable", + telemetry_context=telemetry_context, + original_exception=error, + ) + except ( + json.JSONDecodeError, + UnicodeDecodeError, + PyJWKSetError, + jwt.PyJWKError, + ) as error: + _reject_jwks_dependency( + failure_reason="jwks_invalid", + telemetry_context=telemetry_context, + original_exception=error, + ) + except PyJWKClientError as error: + _reject_jwks_dependency( + failure_reason="jwks_invalid", + telemetry_context=telemetry_context, + original_exception=error, + ) + + if key is None: + _reject_client_jwt( + failure_reason="jwt_unknown_key_id", + telemetry_context=telemetry_context, ) - except jwt.PyJWKSetError as exc: - logger.error(f"Invalid JWKS format: {exc}") - raise AuthException(internal_message=f"Invalid JWKS format: {exc}") + + 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() + 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 +171,188 @@ 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 _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, +) -> _DashboardJWTTelemetry: + is_key_id_present = key_id is not None and bool(key_id.strip()) + jwt_algorithm: str | None = None + if isinstance(algorithm, str) and algorithm in JWT_ALGORITHMS: + jwt_algorithm = algorithm + jwt_kid: str | None = None + if is_key_id_present and key_id is not None: + 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: + sanitized_key_id = JWT_KEY_ID_UNSAFE_PATTERN.sub("_", key_id) + return sanitized_key_id[:JWT_KEY_ID_MAX_LENGTH] + + +def _reject_client_jwt( + *, + failure_reason: JWTFailureReason, + telemetry_context: _DashboardJWTTelemetry, +) -> 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: _DashboardJWTTelemetry, + 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: _DashboardJWTTelemetry, + is_jwks_dependency_failure: bool, + original_exception: Exception | None = None, +) -> None: + log_data: dict[str, object] = { + **telemetry_context.to_log_data(), + "failure_reason": failure_reason, + } + 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: if value == READ_ONLY_PERMISSION: return READ_ONLY_PERMISSION 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_dashboard_jwt_authentication_contract.py b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py new file mode 100644 index 000000000..5d38d3a5a --- /dev/null +++ b/apps/api/tests/contract/test_dashboard_jwt_authentication_contract.py @@ -0,0 +1,652 @@ +from __future__ import annotations + +import base64 +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 pathlib import Path +from typing import 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 tests.support.import_environment import ( + configure_import_environment, + ensure_import_paths, +) +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, +) + +configure_import_environment() +ensure_import_paths() + + +class _LoguruMessage(Protocol): + @property + def record(self) -> Mapping[str, object]: + raise NotImplementedError + + +@dataclass(frozen=True) +class _CapturedAuthLog: + level: str + event: str + message: str + extra: Mapping[str, object] + exception_type: str | None + exception_message: str | None + + +@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: 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, + ) + ) + + +@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 _prepare_api_app_imports() -> None: + 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) + + 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) + + +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() -> BaseException: + from shared.core.exceptions.domain_exceptions import AuthException + + return AuthException() + + +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() + + @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 "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(".") + 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 = _serialize_auth_log(auth_log) + 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 _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( + { + "id": "contract-dashboard-user", + "exp": datetime.now(timezone.utc) + timedelta(minutes=5), + }, + "contract-secret-with-at-least-32-bytes", + algorithm="HS256", + ) + + +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, +) -> None: + token = _create_token_without_key_id() + + with _capture_auth_logs() as log_capture: + with _serve_jwks() as jwks_server: + _use_dashboard_endpoint(monkeypatch, 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_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 + 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: + _use_dashboard_endpoint(monkeypatch, 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_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 + assert "jwt_kid" not in auth_log.extra + _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="Authentication required", + token=token, + ) + assert jwks_server.state.request_count == 0 + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + _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 + 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, +) -> 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")]} + ) + _use_dashboard_endpoint(monkeypatch, 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_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 + 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") + 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(jwks_body, status_code=503) + _use_dashboard_endpoint(monkeypatch, 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_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( + "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) + _use_dashboard_endpoint(monkeypatch, 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_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 +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)]} + ) + _use_dashboard_endpoint(monkeypatch, 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)]} + ) + _use_dashboard_endpoint(monkeypatch, 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_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 + _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)]} + ) + _use_dashboard_endpoint(monkeypatch, 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_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 + _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)]} + ) + _use_dashboard_endpoint(monkeypatch, 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_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 + _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)]} + ) + _use_dashboard_endpoint(monkeypatch, jwks_server.endpoint) + response = await _request_with_token(token) + + _assert_unauthenticated_response( + response, + expected_message="Authentication required", + token=token, + ) + + assert len(log_capture.records) == 1 + auth_log = log_capture.records[0] + _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 + assert "jwt_algorithm" not in auth_log.extra + _assert_log_excludes_token(auth_log, token=token) + + +def test_logfire_exception_callback_downgrades_auth_exceptions_by_status() -> None: + helper = _FakeLogfireExceptionHelper(exception=_create_auth_exception()) + + _downgrade_logfire_exception(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/contract/test_self_hosted_telemetry_contract.py b/apps/api/tests/contract/test_self_hosted_telemetry_contract.py index 54152c258..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,11 +26,21 @@ stop_self_hosted_aggregate_telemetry, ) from shared.services.telemetry.events import ( + 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" @@ -75,7 +85,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, @@ -94,14 +104,180 @@ 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", + "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()) +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( + "oss_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( + "oss_document_type_aggregate", + { + "document_type": "pdf", + "jobs_created_24h": 1, + "source_file_name": "private.pdf", + "email": "user@example.com", + }, + ) + client_properties = sanitize_event_properties( + "oss_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 == "oss_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 == "oss_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( + "oss_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: @@ -132,7 +308,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, @@ -184,8 +360,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 @@ -198,7 +374,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, @@ -243,7 +419,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, @@ -255,7 +431,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" ) @@ -278,7 +454,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, @@ -292,7 +468,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, @@ -339,7 +515,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}", }, @@ -347,9 +523,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 @@ -365,7 +541,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", }, @@ -373,7 +549,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", }, @@ -393,7 +569,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", }, @@ -401,7 +577,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", }, @@ -428,6 +604,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 +700,5 @@ def _build_config( environment="production", app_env="production", service_name="knowhere-api", + schema_version=SCHEMA_VERSION, ) 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/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/docs/adr/0004-anonymous-self-hosted-telemetry.md b/docs/adr/0004-anonymous-self-hosted-telemetry.md new file mode 100644 index 000000000..81a71e087 --- /dev/null +++ b/docs/adr/0004-anonymous-self-hosted-telemetry.md @@ -0,0 +1,111 @@ +# 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 + +Event names use the `oss_` prefix (not `self_hosted_`). Catalog: + +| Event | Role | +| --- | --- | +| `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 + +- `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|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` + +### 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/core/exceptions/knowhere_exception.py b/packages/shared-python/shared/core/exceptions/knowhere_exception.py index 9158be63f..424431a0b 100644 --- a/packages/shared-python/shared/core/exceptions/knowhere_exception.py +++ b/packages/shared-python/shared/core/exceptions/knowhere_exception.py @@ -59,13 +59,14 @@ raise KnowhereException(code=ErrorCode.INVALID_ARGUMENT, ...) """ -from typing import Any, Dict, Optional +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." +LogErrorCategory = Literal["client", "system"] class KnowhereException(Exception): @@ -215,13 +216,10 @@ 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] = { "error_code": self.code.value, "http_status": self.http_status_code, - "error_category": 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, @@ -324,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/models/schemas/job.py b/packages/shared-python/shared/models/schemas/job.py index 4326ba439..2a948f792 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,17 @@ 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. " + "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." + ), + ) + 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..4efb808b3 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] = 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: + 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 + 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..cdcc3d4bd --- /dev/null +++ b/packages/shared-python/shared/models/schemas/llm_config.py @@ -0,0 +1,174 @@ +"""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 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):: + + {"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:: + + { + "text": {"api_key": "...", "model": "...", "base_url": "..."}, + "vision": {"api_key": "...", "model": "...", "base_url": "..."} + } + + Semantics: + - 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 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)", + ) + vision: Optional[LLMProviderConfig] = Field( + None, + description="Vision / VLM credentials (replaces root for vision channel)", + ) + + @model_validator(mode="after") + def _validate_shape(self) -> "LLMConfig": + 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 with api_key/base_url requires model and/or models" + ) + + 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 _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=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_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_for("vision") + ) + + 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 { + "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), + } + + +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..bc5671ecd --- /dev/null +++ b/packages/shared-python/shared/services/ai/llm_overrides.py @@ -0,0 +1,160 @@ +"""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 + except ImportError: + # 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 + + +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, ) 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..80dfe9ec7 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]] = [ + ( + "oss_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( + ( + "oss_usage_aggregate", + await _collect_usage_aggregate(session, config, window_seconds), + ) + ) + captures.append( + ( + "oss_retrieval_aggregate", + await _collect_retrieval_aggregate( session, config, window_seconds, ), - } + ) ) - return event_properties + captures.append( + ( + "oss_worker_aggregate", + await _collect_worker_aggregate(session, config, window_seconds), + ) + ) + captures.append( + ( + "oss_provider_aggregate", + await _collect_provider_aggregate(session, config, window_seconds), + ) + ) + for properties in await _collect_document_type_aggregates( + session, + config, + window_seconds, + ): + captures.append(("oss_document_type_aggregate", properties)) + for properties in await _collect_client_aggregates( + session, + config, + window_seconds, + ): + captures.append(("oss_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 { + "oss_document_type_aggregate", + "oss_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..827b10b2b 100644 --- a/packages/shared-python/shared/services/telemetry/events.py +++ b/packages/shared-python/shared/services/telemetry/events.py @@ -4,6 +4,7 @@ import os from collections.abc import Mapping +from pathlib import PurePosixPath from typing import TypeAlias, cast from .config import TelemetryRuntimeConfig @@ -11,6 +12,46 @@ 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", + "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", @@ -47,8 +88,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", @@ -56,8 +97,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", @@ -66,8 +107,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", @@ -75,7 +125,7 @@ "total_users", } ), - "self_hosted_retrieval_aggregate": _AGGREGATE_PROPERTY_NAMES + "oss_retrieval_aggregate": _AGGREGATE_PROPERTY_NAMES | frozenset( { "retrieval_cache_hits_24h", @@ -89,7 +139,7 @@ "retrieval_tokens_24h", } ), - "self_hosted_worker_aggregate": _AGGREGATE_PROPERTY_NAMES + "oss_worker_aggregate": _AGGREGATE_PROPERTY_NAMES | frozenset( { "job_duration_avg_seconds_24h", @@ -101,7 +151,7 @@ "jobs_waiting_file", } ), - "self_hosted_api_aggregate": _AGGREGATE_PROPERTY_NAMES + "oss_api_aggregate": _AGGREGATE_PROPERTY_NAMES | frozenset( { "api_latency_avg_ms", @@ -113,7 +163,7 @@ "api_requests_total", } ), - "self_hosted_provider_aggregate": _AGGREGATE_PROPERTY_NAMES + "oss_provider_aggregate": _AGGREGATE_PROPERTY_NAMES | frozenset( { "parse_agent_errors_24h", @@ -128,6 +178,27 @@ "webhook_delivery_failures_24h", } ), + "oss_document_type_aggregate": _AGGREGATE_PROPERTY_NAMES + | frozenset( + { + "document_type", + "jobs_created_24h", + "jobs_done_24h", + "jobs_failed_24h", + "pages_processed_24h", + "success_rate_24h", + } + ), + "oss_client_aggregate": _AGGREGATE_PROPERTY_NAMES + | frozenset( + { + "created_by_client", + "jobs_created_24h", + "jobs_done_24h", + "jobs_failed_24h", + "success_rate_24h", + } + ), } @@ -136,6 +207,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..5a63fab4a 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( + "oss_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, @@ -52,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, @@ -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("oss_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 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..70d7eb6d6 --- /dev/null +++ b/packages/shared-python/shared/tests/test_llm_config.py @@ -0,0 +1,150 @@ +"""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, LLMModelsConfig, LLMProviderConfig + + +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_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().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" + 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_two_different_endpoints() -> None: + cfg = LLMConfig( + 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", 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_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="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() + + +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["api_key"] != "sk-root" + assert dump["model"] == "gpt-4o" + assert dump["vision"]["api_key"] != "sk-vision" + assert dump["text"] is None diff --git a/scripts/smoke_byok.sh b/scripts/smoke_byok.sh new file mode 100755 index 000000000..104848f78 --- /dev/null +++ b/scripts/smoke_byok.sh @@ -0,0 +1,153 @@ +#!/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 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-flat", + "llm_config":{ + "api_key":"sk-smoke-test-key-not-real", + "model":"gpt-4o", + "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 "== 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 '{ + "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 '{ + "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."