From 3fd786832e8169e0c5b23e11a14e43535aac1dbb Mon Sep 17 00:00:00 2001 From: suguanYang Date: Fri, 1 May 2026 06:39:01 +0000 Subject: [PATCH 01/32] fix: address CodeQL security findings --- README.md | 4 + .../f6a7b8c9d0e1_key_api_key_hashes.py | 31 ++++++ apps/api/app/api/v1/health.py | 16 +-- apps/api/app/api/v1/routes/s3_events.py | 97 ++++++++++++++++--- apps/api/app/core/dependencies.py | 4 +- apps/api/app/services/auth/api_key_service.py | 10 +- .../guest/guest_registration_service.py | 11 +-- .../app/services/rate_limit/dependencies.py | 6 +- apps/api/scripts/bootstrap_local_dev.py | 11 +-- apps/api/scripts/init_user.py | 34 ++++++- .../scripts/local_dev_bootstrap_service.py | 17 +++- .../contract/test_database_health_contract.py | 37 ++++--- .../contract/test_job_creation_contract.py | 71 +++++++++++++- .../tests/contract/test_s3_event_contract.py | 59 +++++++++++ apps/api/tests/support/contract_database.py | 7 +- deploy/local-dev/README.md | 2 +- .../shared-python/shared/core/database.py | 4 +- .../shared/models/database/api_key.py | 3 + .../shared/testing/contract_runtime.py | 2 +- .../shared/utils/api_key_hashing.py | 13 +++ .../shared/utils/url_file_type.py | 68 ++++++++++--- .../shared/utils/url_security.py | 65 +++++++++++-- 22 files changed, 472 insertions(+), 100 deletions(-) create mode 100644 apps/api/alembic/versions/f6a7b8c9d0e1_key_api_key_hashes.py create mode 100644 packages/shared-python/shared/utils/api_key_hashing.py diff --git a/README.md b/README.md index 16e463de2..3057d4143 100644 --- a/README.md +++ b/README.md @@ -83,6 +83,10 @@ uv run --python 3.11 python -m alembic upgrade heads uv run --python 3.11 python scripts/init_user.py --email you@example.com ``` +Pass `--api-key-output-file ./standalone-api-key.txt` if you need the generated +plaintext key written to a local `0600` file. The default console output only +reports that the credential was created. + If you plan to use the dashboard, start the combined self-hosted stack and register through the dashboard instead of using `scripts/init_user.py`. diff --git a/apps/api/alembic/versions/f6a7b8c9d0e1_key_api_key_hashes.py b/apps/api/alembic/versions/f6a7b8c9d0e1_key_api_key_hashes.py new file mode 100644 index 000000000..cb7688d90 --- /dev/null +++ b/apps/api/alembic/versions/f6a7b8c9d0e1_key_api_key_hashes.py @@ -0,0 +1,31 @@ +"""key api key hashes + +Revision ID: f6a7b8c9d0e1 +Revises: e5f6a7b8c9d0 +Create Date: 2026-05-01 05:55:00.000000 + +""" + +from typing import Sequence, Union + +import sqlalchemy as sa +from alembic import op + + +revision: str = "f6a7b8c9d0e1" +down_revision: Union[str, Sequence[str], None] = "e5f6a7b8c9d0" +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + op.execute("UPDATE api_keys SET is_active = false") + op.add_column( + "api_keys", + sa.Column("hash_version", sa.String(length=16), nullable=False, server_default="hmac-v1"), + ) + op.alter_column("api_keys", "hash_version", server_default=None) + + +def downgrade() -> None: + op.drop_column("api_keys", "hash_version") diff --git a/apps/api/app/api/v1/health.py b/apps/api/app/api/v1/health.py index f3fc2b2dd..f01472202 100644 --- a/apps/api/app/api/v1/health.py +++ b/apps/api/app/api/v1/health.py @@ -6,7 +6,6 @@ from shared.core.database import ( get_database_health, - get_database_info, get_database_performance, prewarm_connection_pool, ) @@ -17,13 +16,14 @@ @router.get("/database/health") async def check_database_health(): """Check database health status.""" - return await get_database_health() - - -@router.get("/database/info") -async def get_database_information(): - """Return database connection information.""" - return await get_database_info() + health = await get_database_health() + if "error" in health: + return { + "status": health.get("status", "unhealthy"), + "error": "Database health check failed", + "last_check": health.get("last_check"), + } + return health @router.get("/database/performance") diff --git a/apps/api/app/api/v1/routes/s3_events.py b/apps/api/app/api/v1/routes/s3_events.py index a65990179..9b0897ca0 100644 --- a/apps/api/app/api/v1/routes/s3_events.py +++ b/apps/api/app/api/v1/routes/s3_events.py @@ -5,9 +5,11 @@ import base64 import json import os +import socket from typing import Any, Dict import aiohttp +from aiohttp.abc import AbstractResolver from app.repositories.job_repository import JobRepository from app.services.knowledge.kb_orchestrator import KBOrchestrator from app.services.state_machine import JobStateMachine @@ -19,9 +21,13 @@ from shared.core.state_machine.states import JobStatus from shared.models.schemas.oss_event import OSSEvent from shared.models.schemas.s3_event import S3Event +from shared.services.webhook.validator import validate_webhook_url_async +from shared.utils.url_security import SafePublicHTTPURL router = APIRouter(tags=["Internal"]) +SNS_SUBSCRIPTION_TIMEOUT_SECONDS = 10 + def verify_sns_signature(request_body: bytes, signature: str, message: str) -> bool: """ @@ -207,22 +213,7 @@ async def handle_sns_event(body: bytes): if subscribe_url: logger.info(f"SNS subscription confirmation URL: {subscribe_url}") # Visit the URL to confirm the subscription. - try: - async with aiohttp.ClientSession() as session: - async with session.get(subscribe_url) as response: - if response.status == 200: - logger.info("SNS subscription confirmed successfully") - return {"message": "SNS subscription confirmed"} - else: - logger.error( - f"SNS subscription confirmation failed, status={response.status}" - ) - return { - "message": "SNS subscription confirmation failed" - } - except Exception as e: - logger.error(f"Failed to reach the SNS confirmation URL: {e}") - return {"message": "SNS subscription confirmation failed"} + return await confirm_sns_subscription(subscribe_url) else: logger.warning( "SNS subscription confirmation did not include SubscribeURL" @@ -274,6 +265,80 @@ async def handle_sns_event(body: bytes): raise +async def confirm_sns_subscription(subscribe_url: str) -> dict[str, str]: + """Confirm an SNS subscription after SSRF validation and IP pinning.""" + validation = await validate_webhook_url_async(subscribe_url) + if not validation.is_valid: + logger.warning( + f"SNS subscription confirmation URL failed validation: {validation.error_message}" + ) + return {"message": "SNS subscription confirmation failed"} + + if not validation.validated_ip: + logger.warning("SNS subscription confirmation URL validation returned no IP") + return {"message": "SNS subscription confirmation failed"} + + try: + validated_subscribe_url = SafePublicHTTPURL(subscribe_url) + connector = aiohttp.TCPConnector( + resolver=_PinnedSNSResolver(validation.validated_ip), + ) + timeout = aiohttp.ClientTimeout(total=SNS_SUBSCRIPTION_TIMEOUT_SECONDS) + async with aiohttp.ClientSession( + connector=connector, + timeout=timeout, + ) as session: + async with session.get( + validated_subscribe_url, + allow_redirects=False, + ) as response: + if response.status == 200: + logger.info("SNS subscription confirmed successfully") + return {"message": "SNS subscription confirmed"} + + if 300 <= response.status < 400: + logger.warning( + f"SNS subscription confirmation redirect blocked, status={response.status}" + ) + else: + logger.error( + f"SNS subscription confirmation failed, status={response.status}" + ) + return {"message": "SNS subscription confirmation failed"} + except Exception as e: + logger.error(f"Failed to reach the SNS confirmation URL: {e}") + return {"message": "SNS subscription confirmation failed"} + + +class _PinnedSNSResolver(AbstractResolver): + """Resolver that pins SNS confirmation to a pre-validated public IP.""" + + def __init__(self, pinned_ip: str) -> None: + self.pinned_ip = pinned_ip + + async def resolve( + self, + host: str, + port: int = 0, + family: int = socket.AF_INET, + ) -> list[dict[str, Any]]: + parsed_ip = self.pinned_ip + pinned_family = socket.AF_INET6 if ":" in parsed_ip else socket.AF_INET + return [ + { + "hostname": host, + "host": parsed_ip, + "port": port, + "family": pinned_family, + "proto": 0, + "flags": socket.AI_NUMERICHOST, + } + ] + + async def close(self) -> None: + pass + + async def handle_minio_event(body: bytes, auth_token: str): """ Handle a MinIO webhook event. diff --git a/apps/api/app/core/dependencies.py b/apps/api/app/core/dependencies.py index ac349e34f..a234bb92e 100644 --- a/apps/api/app/core/dependencies.py +++ b/apps/api/app/core/dependencies.py @@ -1,4 +1,3 @@ -import hashlib import threading from datetime import timedelta from fnmatch import fnmatch @@ -18,6 +17,7 @@ AuthException, PermissionDeniedException, ) +from shared.utils.api_key_hashing import hash_api_key # Standard JWKS endpoint path (fixed, following OpenID Connect convention) JWKS_ENDPOINT_PATH = "/api/auth/jwks" @@ -186,7 +186,7 @@ async def get_current_user_id( # Mode 1: API Key verification (for external clients) if token.startswith("sk_"): # Check identity cache first — skip DB on cache hit - api_key_hash = hashlib.sha256(token.encode()).hexdigest() + api_key_hash = hash_api_key(token) try: cached = await identity_cache.get_cached_identity( redis_pool_manager.get_redis_service(), diff --git a/apps/api/app/services/auth/api_key_service.py b/apps/api/app/services/auth/api_key_service.py index 920751a00..c9f50d08c 100644 --- a/apps/api/app/services/auth/api_key_service.py +++ b/apps/api/app/services/auth/api_key_service.py @@ -1,7 +1,6 @@ """API key management service.""" import asyncio -import hashlib import uuid from dataclasses import dataclass from datetime import datetime @@ -23,6 +22,7 @@ ) from shared.models.database.api_key import APIKey from shared.models.database.user_balance import UserBalance +from shared.utils.api_key_hashing import hash_api_key _DEFAULT_USER_TIER: str = "free" @@ -85,7 +85,7 @@ async def create_api_key( # 3. Generate a secure API key (sk_ + a 32-char UUID without hyphens). api_key = f"sk_{str(uuid.uuid4()).replace('-', '')}" - key_hash = hashlib.sha256(api_key.encode()).hexdigest() + key_hash = hash_api_key(api_key) key_mask = self._mask_api_key(api_key) # 4. Store it in the database. @@ -116,7 +116,7 @@ async def validate_api_key_identity( api_key: str, ) -> Optional[APIKeyIdentity]: """Validate API key and return the authenticated identity.""" - key_hash = hashlib.sha256(api_key.encode()).hexdigest() + key_hash = hash_api_key(api_key) api_key_record = await self.repository.get_by_key_hash(session, key_hash) if not api_key_record or not api_key_record.is_valid(): @@ -236,7 +236,7 @@ async def regenerate_api_key( # 2. Generate a new API key (sk_ + a 32-char UUID without hyphens). new_api_key = f"sk_{str(uuid.uuid4()).replace('-', '')}" - new_key_hash = hashlib.sha256(new_api_key.encode()).hexdigest() + new_key_hash = hash_api_key(new_api_key) new_key_mask = self._mask_api_key(new_api_key) # 3. Update the database record. @@ -267,7 +267,7 @@ async def check_module_permission( self, session: AsyncSession, api_key: str, module: str ) -> bool: """Check whether an API key can access the requested module.""" - key_hash = hashlib.sha256(api_key.encode()).hexdigest() + key_hash = hash_api_key(api_key) api_key_record = await self.repository.get_by_key_hash(session, key_hash) if not api_key_record or not api_key_record.is_valid(): diff --git a/apps/api/app/services/guest/guest_registration_service.py b/apps/api/app/services/guest/guest_registration_service.py index 351663643..f7e5fcafe 100644 --- a/apps/api/app/services/guest/guest_registration_service.py +++ b/apps/api/app/services/guest/guest_registration_service.py @@ -1,6 +1,7 @@ """Guest registration business logic.""" import hashlib +import uuid from datetime import datetime from typing import NoReturn from uuid import uuid4 @@ -25,6 +26,7 @@ GuestRegisterResponse, ) from shared.services.billing.credits_service import CreditsService +from shared.utils.api_key_hashing import hash_api_key _GUEST_TIER: str = "guest" _GUEST_KEY_NAME_PREFIX: str = "guest-device" @@ -153,13 +155,10 @@ async def _create_api_key_without_commit( This avoids the internal commit inside APIKeyService.create_api_key() which would make the key durable before the device row is inserted. """ - import hashlib - import uuid - from shared.models.database.api_key import APIKey api_key = f"sk_{str(uuid.uuid4()).replace('-', '')}" - key_hash = hashlib.sha256(api_key.encode()).hexdigest() + key_hash = hash_api_key(api_key) key_mask = self._api_key_service._mask_api_key(api_key) api_key_record = APIKey( @@ -264,13 +263,11 @@ def _raise_existing_device_conflict(cls, device_id: str) -> NoReturn: @staticmethod async def _resolve_api_key_id(session: AsyncSession, api_key: str) -> str | None: """Resolve the DB id for a just-created API key by its hash.""" - import hashlib - from sqlalchemy import select from shared.models.database.api_key import APIKey - key_hash = hashlib.sha256(api_key.encode()).hexdigest() + key_hash = hash_api_key(api_key) result = await session.execute( select(APIKey.id).where(APIKey.key_hash == key_hash).limit(1) ) diff --git a/apps/api/app/services/rate_limit/dependencies.py b/apps/api/app/services/rate_limit/dependencies.py index c34b1c10b..c70f6d8c5 100644 --- a/apps/api/app/services/rate_limit/dependencies.py +++ b/apps/api/app/services/rate_limit/dependencies.py @@ -14,7 +14,6 @@ only when billing is enabled. """ -import hashlib import math from datetime import datetime, timezone from typing import AsyncGenerator @@ -42,6 +41,7 @@ from shared.core.logging import log_context from shared.core.state_machine.states import JobStatus from shared.models.database.api_key import APIKey +from shared.utils.api_key_hashing import hash_api_key from shared.models.database.job import Job from shared.models.database.user_balance import UserBalance @@ -175,7 +175,7 @@ async def with_current_user( api_key_hash = None is_api_key_auth = isinstance(token, str) and token.startswith("sk_") if token is not None and is_api_key_auth: - api_key_hash = hashlib.sha256(token.encode()).hexdigest() + api_key_hash = hash_api_key(token) if is_api_key_auth and api_key_hash: try: ttl_seconds = await _resolve_apikey_cache_ttl_seconds(api_key_hash) @@ -197,7 +197,7 @@ async def with_current_user( api_key_hash = None is_api_key_auth = isinstance(token, str) and token.startswith("sk_") if token is not None and is_api_key_auth: - api_key_hash = hashlib.sha256(token.encode()).hexdigest() + api_key_hash = hash_api_key(token) cache_key: str = ( identity_cache._apikey_key(api_key_hash) if is_api_key_auth and api_key_hash diff --git a/apps/api/scripts/bootstrap_local_dev.py b/apps/api/scripts/bootstrap_local_dev.py index 779dee503..e33e895a9 100644 --- a/apps/api/scripts/bootstrap_local_dev.py +++ b/apps/api/scripts/bootstrap_local_dev.py @@ -45,12 +45,11 @@ async def _run(mode: str) -> int: def _print_profile() -> None: - profile = LocalDevelopmentBootstrapService.get_local_developer_profile() - print(f"user_id={profile['user_id']}") - print(f"name={profile['name']}") - print(f"email={profile['email']}") - print(f"tier={profile['tier']}") - print(f"api_key={profile['api_key']}") + print("user_id=local-dev-user") + print("name=Local Development User") + print("email=local-dev-user@knowhere.local") + print("tier=tier_5") + print("local_developer_key_seeded=true") def main() -> int: diff --git a/apps/api/scripts/init_user.py b/apps/api/scripts/init_user.py index 2c156ed0e..0ae8fa904 100644 --- a/apps/api/scripts/init_user.py +++ b/apps/api/scripts/init_user.py @@ -2,10 +2,10 @@ import argparse import asyncio -import hashlib import os import secrets import sys +from pathlib import Path from uuid import uuid4 from sqlalchemy import select @@ -20,6 +20,7 @@ from shared.models.database.user import User from shared.models.database.user_balance import UserBalance from shared.services.auth.user_table_bootstrap import ensure_better_auth_user_table +from shared.utils.api_key_hashing import hash_api_key _DEFAULT_API_KEY_NAME: str = "standalone-api-key" _DEFAULT_USER_TIER: str = "free" @@ -43,6 +44,11 @@ def _build_parser() -> argparse.ArgumentParser: default=_DEFAULT_USER_TIER, help="Compatibility user tier to store in user_balances.", ) + parser.add_argument( + "--api-key-output-file", + default="", + help="Optional file path for the generated API key. The file is created with 0600 permissions.", + ) return parser @@ -134,6 +140,19 @@ def _mask_api_key(api_key: str) -> str: return api_key[:8] + "•" * (len(api_key) - 12) + api_key[-4:] +def _write_api_key_file(path_value: str, api_key: str) -> Path: + output_path = Path(path_value).expanduser() + output_path.parent.mkdir(parents=True, exist_ok=True) + file_descriptor = os.open( + output_path, + os.O_WRONLY | os.O_CREAT | os.O_TRUNC, + 0o600, + ) + with os.fdopen(file_descriptor, "w", encoding="utf-8") as output_file: + output_file.write(f"{api_key}\n") + return output_path + + async def _create_api_key( session: AsyncSession, *, @@ -144,7 +163,7 @@ async def _create_api_key( session.add( APIKey( user_id=user_id, - key_hash=hashlib.sha256(api_key.encode()).hexdigest(), + key_hash=hash_api_key(api_key), key_mask=_mask_api_key(api_key), name=key_name, enabled_modules=["all"], @@ -179,8 +198,15 @@ async def _run(args: argparse.Namespace) -> int: print(f"user_id={user.id}") print(f"email={user.email}") - print(f"api_key_name={key_name}") - print(f"api_key={api_key}") + print("credential_name_created=true") + print("api_key_created=true") + print("credential_hidden=true") + output_path_value = str(args.api_key_output_file).strip() + if output_path_value: + output_path = _write_api_key_file(output_path_value, api_key) + print(f"credential_output_file={output_path}") + else: + print("credential_output_file=") return 0 diff --git a/apps/api/scripts/local_dev_bootstrap_service.py b/apps/api/scripts/local_dev_bootstrap_service.py index 758ecaab6..d4ac0217c 100644 --- a/apps/api/scripts/local_dev_bootstrap_service.py +++ b/apps/api/scripts/local_dev_bootstrap_service.py @@ -1,6 +1,5 @@ from __future__ import annotations -import hashlib from datetime import datetime, timezone from sqlalchemy.ext.asyncio import AsyncSession @@ -13,6 +12,7 @@ from shared.models.database.user import User from shared.models.database.user_balance import UserBalance from shared.services.auth.user_table_bootstrap import ensure_better_auth_user_table +from shared.utils.api_key_hashing import hash_api_key class LocalDevelopmentBootstrapService: @@ -52,16 +52,23 @@ async def seed_local_developer(self, session: AsyncSession) -> None: @classmethod def get_local_developer_profile(cls) -> dict[str, str | int]: - """Expose deterministic local developer credentials for local tooling.""" - return { + """Expose deterministic local developer profile details for local tooling.""" + profile: dict[str, str | int] = { "user_id": cls.LOCAL_DEV_USER_ID, "name": cls.LOCAL_DEV_USER_NAME, "email": cls.LOCAL_DEV_USER_EMAIL, "tier": cls.LOCAL_DEV_TIER, - "api_key": cls.LOCAL_DEV_API_KEY, "credits_balance": cls.LOCAL_DEV_CREDITS_BALANCE, "lifetime_billing_micro": cls.LOCAL_DEV_LIFETIME_BILLING_MICRO, } + return profile + + @classmethod + def get_local_developer_auth_profile(cls) -> dict[str, str | int]: + """Expose deterministic local developer auth details for contract tests.""" + auth_profile = cls.get_local_developer_profile() + auth_profile["api_key"] = cls.LOCAL_DEV_API_KEY + return auth_profile async def _upsert_user(self, session: AsyncSession) -> None: user = await session.get(User, self.LOCAL_DEV_USER_ID) @@ -148,7 +155,7 @@ async def _upsert_credits_transaction(self, session: AsyncSession) -> None: async def _upsert_api_key(self, session: AsyncSession) -> None: api_key = await session.get(APIKey, self.LOCAL_DEV_API_KEY_ID) - key_hash = hashlib.sha256(self.LOCAL_DEV_API_KEY.encode()).hexdigest() + key_hash = hash_api_key(self.LOCAL_DEV_API_KEY) key_mask = self._mask_api_key(self.LOCAL_DEV_API_KEY) if api_key is None: diff --git a/apps/api/tests/contract/test_database_health_contract.py b/apps/api/tests/contract/test_database_health_contract.py index 645620964..2456ec298 100644 --- a/apps/api/tests/contract/test_database_health_contract.py +++ b/apps/api/tests/contract/test_database_health_contract.py @@ -4,6 +4,7 @@ import pytest from httpx import AsyncClient +from pytest import MonkeyPatch @pytest.mark.asyncio @@ -31,26 +32,32 @@ async def test_should_return_the_database_health_payload_shape( @pytest.mark.asyncio -async def test_should_return_the_database_info_payload_shape( +async def test_should_sanitize_database_health_errors_before_returning_them( api_client_factory: Callable[[], AbstractAsyncContextManager[AsyncClient]], + monkeypatch: MonkeyPatch, ) -> None: - async with api_client_factory() as api_client: - response = await api_client.get("/api/v1/health/database/info") + async def _get_database_health() -> dict[str, object]: + return { + "status": "unhealthy", + "error": "Database health check failed", + "last_check": None, + } - assert response.status_code == 200 + async with api_client_factory() as api_client: + from app.api.v1 import health as health_route_module - response_json = cast(dict[str, object], response.json()) - pool_status = cast(dict[str, object], response_json["pool_status"]) + monkeypatch.setattr( + health_route_module, + "get_database_health", + _get_database_health, + ) + response = await api_client.get("/api/v1/health/database/health") - assert response_json["version"] - assert isinstance(response_json["active_connections"], int) - assert response_json["database_size"] - assert set(pool_status) == { - "size", - "checked_in", - "checked_out", - "overflow", - "invalid", + assert response.status_code == 200 + assert response.json() == { + "status": "unhealthy", + "error": "Database health check failed", + "last_check": None, } diff --git a/apps/api/tests/contract/test_job_creation_contract.py b/apps/api/tests/contract/test_job_creation_contract.py index 71a3ad4ee..69f9ddce4 100644 --- a/apps/api/tests/contract/test_job_creation_contract.py +++ b/apps/api/tests/contract/test_job_creation_contract.py @@ -532,8 +532,9 @@ async def test_should_create_a_waiting_file_job_for_a_url_source_and_enqueue_the scheduled_tasks: list[dict[str, object]] = [] class _FakeHeadResponse: - def __init__(self, content_type: str) -> None: + def __init__(self, content_type: str, status_code: int = 200) -> None: self.headers: dict[str, str] = {"content-type": content_type} + self.status_code = status_code class _FakeAsyncHttpClient: async def head( @@ -543,7 +544,7 @@ async def head( follow_redirects: bool = True, ) -> _FakeHeadResponse: requested_urls.append(url) - assert follow_redirects is True + assert follow_redirects is False return _FakeHeadResponse("application/pdf") class _FakeCeleryTask: @@ -669,6 +670,8 @@ async def test_should_reject_url_source_when_url_resolves_to_private_network( def resolve_private_address( host: str, port: int | None, + *args: object, + **kwargs: object, ) -> list[tuple[socket.AddressFamily, socket.SocketKind, int, str, tuple[str, int]]]: return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("127.0.0.1", 0))] @@ -697,6 +700,70 @@ def resolve_private_address( assert await _count_jobs() == 0 +@pytest.mark.asyncio +async def test_should_reject_a_url_source_when_file_type_detection_redirects_to_a_private_host( + monkeypatch: MonkeyPatch, + developer_api_client_factory: Callable[ + [], AbstractAsyncContextManager[AsyncClient] + ], +) -> None: + payload: dict[str, str] = { + "namespace": "contract-jobs", + "source_type": "url", + "source_url": "https://example.com/contracts/knowhere-upload", + "data_id": "contract-job-url-private-redirect", + } + requested_urls: list[str] = [] + + class _FakeHeadResponse: + status_code = 302 + headers: dict[str, str] = { + "location": "http://127.0.0.1/internal-metadata.pdf", + } + + class _FakeAsyncHttpClient: + async def head( + self, + url: str, + *, + follow_redirects: bool = True, + ) -> _FakeHeadResponse: + requested_urls.append(url) + assert follow_redirects is False + return _FakeHeadResponse() + + import shared.utils.http_clients as http_clients_module + + monkeypatch.setattr( + http_clients_module, + "get_async_client", + lambda: _FakeAsyncHttpClient(), + ) + + async with developer_api_client_factory() as api_client: + response = await api_client.post("/api/v1/jobs", json=payload) + + assert response.status_code == 400 + assert response.headers["x-request-id"] + + response_json: dict[str, object] = response.json() + error = cast(dict[str, object], response_json["error"]) + details = cast(dict[str, object], error["details"]) + violations = cast(list[dict[str, object]], details["violations"]) + + assert requested_urls == [payload["source_url"]] + assert response_json["success"] is False + assert error["code"] == "INVALID_ARGUMENT" + assert error["message"] == "Invalid URL" + assert violations == [ + { + "field": "source_url", + "description": "URL host is not allowed", + } + ] + assert await _count_jobs() == 0 + + @pytest.mark.asyncio async def test_should_confirm_upload_and_start_processing_for_a_waiting_file_job( monkeypatch: MonkeyPatch, diff --git a/apps/api/tests/contract/test_s3_event_contract.py b/apps/api/tests/contract/test_s3_event_contract.py index c2cd0899b..5b07b2fb7 100644 --- a/apps/api/tests/contract/test_s3_event_contract.py +++ b/apps/api/tests/contract/test_s3_event_contract.py @@ -1,5 +1,6 @@ import importlib import json +import socket from collections.abc import Callable from contextlib import AbstractAsyncContextManager from typing import cast @@ -160,6 +161,64 @@ async def start_workflow( assert job_row["status"] == "pending" +@pytest.mark.asyncio +async def test_should_reject_an_sns_subscription_confirmation_url_that_resolves_to_a_private_host( + api_client_factory: Callable[[], AbstractAsyncContextManager[AsyncClient]], + monkeypatch: MonkeyPatch, +) -> None: + contacted_urls: list[str] = [] + + def resolve_private_address( + host: str, + port: int | None, + *args: object, + **kwargs: object, + ) -> list[tuple[socket.AddressFamily, socket.SocketKind, int, str, tuple[str, int]]]: + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("127.0.0.1", 0))] + + class _UnexpectedSession: + def __init__(self, *args: object, **kwargs: object) -> None: + pass + + async def __aenter__(self) -> "_UnexpectedSession": + return self + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc: BaseException | None, + traceback: object, + ) -> None: + return None + + def get(self, url: str, *args: object, **kwargs: object) -> object: + contacted_urls.append(url) + raise AssertionError("private SNS confirmation URL should not be requested") + + async with api_client_factory() as api_client: + s3_events_module = importlib.import_module("app.api.v1.routes.s3_events") + monkeypatch.setattr(socket, "getaddrinfo", resolve_private_address) + monkeypatch.setattr( + s3_events_module.aiohttp, + "ClientSession", + _UnexpectedSession, + ) + response = await api_client.post( + "/api/v1/internal/s3-events", + content=json.dumps( + { + "Type": "SubscriptionConfirmation", + "SubscribeURL": "https://sns.example.test/confirm", + } + ).encode("utf-8"), + headers={"x-amz-sns-message-type": "SubscriptionConfirmation"}, + ) + + assert response.status_code == 200 + assert response.json() == {"message": "SNS subscription confirmation failed"} + assert contacted_urls == [] + + @pytest.mark.asyncio async def test_should_return_ok_for_a_malformed_event_payload_without_triggering_retries( api_client_factory: Callable[[], AbstractAsyncContextManager[AsyncClient]], diff --git a/apps/api/tests/support/contract_database.py b/apps/api/tests/support/contract_database.py index 681417188..f704ff41c 100644 --- a/apps/api/tests/support/contract_database.py +++ b/apps/api/tests/support/contract_database.py @@ -1,6 +1,5 @@ from __future__ import annotations -import hashlib import json from datetime import datetime, timezone from typing import Any @@ -10,6 +9,7 @@ from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine from shared.testing.contract_runtime import get_contract_database_url +from shared.utils.api_key_hashing import hash_api_key async def _create_contract_engine() -> AsyncEngine: @@ -139,7 +139,7 @@ async def insert_authenticated_user( ) timestamp = _utc_now() - api_key_hash = hashlib.sha256(api_key.encode()).hexdigest() + api_key_hash = hash_api_key(api_key) api_key_id = f"key_{uuid4().hex[:12]}" await cls.execute( @@ -148,6 +148,7 @@ async def insert_authenticated_user( id, user_id, key_hash, + hash_version, key_mask, name, enabled_modules, @@ -157,6 +158,7 @@ async def insert_authenticated_user( :id, :user_id, :key_hash, + :hash_version, :key_mask, :name, CAST(:enabled_modules AS JSON), @@ -168,6 +170,7 @@ async def insert_authenticated_user( "id": api_key_id, "user_id": user_id, "key_hash": api_key_hash, + "hash_version": "hmac-v1", "key_mask": f"{api_key[:8]}...{api_key[-4:]}", "name": f"Contract API Key {user_id}", "enabled_modules": json.dumps(enabled_modules or ["all"]), diff --git a/deploy/local-dev/README.md b/deploy/local-dev/README.md index e481b8033..9b102c66d 100644 --- a/deploy/local-dev/README.md +++ b/deploy/local-dev/README.md @@ -45,7 +45,7 @@ Deterministic local developer account: - `user_id`: `local-dev-user` - `email`: `local-dev-user@knowhere.local` - `tier`: `tier_5` -- `api_key`: `local_dev_demo_key_tier5_full_access` +- `local_developer_key_seeded`: `true` ## Verify the Local API diff --git a/packages/shared-python/shared/core/database.py b/packages/shared-python/shared/core/database.py index c1f950f96..acdf08654 100644 --- a/packages/shared-python/shared/core/database.py +++ b/packages/shared-python/shared/core/database.py @@ -196,7 +196,7 @@ async def check_health(self) -> dict[str, object]: logger.error(f"Database health check failed: {e}") return { "status": "unhealthy", - "error": str(e), + "error": "Database health check failed", "last_check": self.last_check.isoformat() if self.last_check else None, } @@ -254,7 +254,7 @@ async def get_database_info(self) -> dict[str, object]: } except Exception as e: logger.error(f"Failed to get database info: {e}") - return {"error": str(e)} + return {"error": "Database information unavailable"} # Shared health-checker instance. diff --git a/packages/shared-python/shared/models/database/api_key.py b/packages/shared-python/shared/models/database/api_key.py index a70b24f9f..55028c1f6 100644 --- a/packages/shared-python/shared/models/database/api_key.py +++ b/packages/shared-python/shared/models/database/api_key.py @@ -29,6 +29,9 @@ class APIKey(Base): key_hash: Mapped[str] = mapped_column( String(255), nullable=False, index=True ) # Encrypted storage + hash_version: Mapped[str] = mapped_column( + String(16), default="hmac-v1", nullable=False + ) key_mask: Mapped[str] = mapped_column( String(50), nullable=False ) # Masked API Key (for display) diff --git a/packages/shared-python/shared/testing/contract_runtime.py b/packages/shared-python/shared/testing/contract_runtime.py index 8c816ce79..5ea5cfd29 100644 --- a/packages/shared-python/shared/testing/contract_runtime.py +++ b/packages/shared-python/shared/testing/contract_runtime.py @@ -657,7 +657,7 @@ async def seed_contract_developer() -> dict[str, str | int]: finally: await engine.dispose() - return bootstrap_module.LocalDevelopmentBootstrapService.get_local_developer_profile() + return bootstrap_module.LocalDevelopmentBootstrapService.get_local_developer_auth_profile() async def reset_contract_database() -> None: diff --git a/packages/shared-python/shared/utils/api_key_hashing.py b/packages/shared-python/shared/utils/api_key_hashing.py new file mode 100644 index 000000000..46dfe705a --- /dev/null +++ b/packages/shared-python/shared/utils/api_key_hashing.py @@ -0,0 +1,13 @@ +"""API key hashing helpers.""" + +import hmac +from hashlib import sha256 + + +def hash_api_key(api_key: str) -> str: + """Return a deterministic keyed digest for API key lookup.""" + from shared.core.config import settings + + secret_key = settings.SECRET_KEY.encode("utf-8") + api_key_bytes = api_key.encode("utf-8") + return hmac.new(secret_key, api_key_bytes, sha256).hexdigest() diff --git a/packages/shared-python/shared/utils/url_file_type.py b/packages/shared-python/shared/utils/url_file_type.py index 38efeb1ec..0a0793a18 100644 --- a/packages/shared-python/shared/utils/url_file_type.py +++ b/packages/shared-python/shared/utils/url_file_type.py @@ -11,7 +11,13 @@ from loguru import logger from shared.core.config import settings -from shared.utils.url_security import validate_public_http_url +from shared.core.exceptions.domain_exceptions import ValidationException +from shared.utils.url_security import ( + MAX_SAFE_REDIRECTS, + SafePublicHTTPURL, + get_safe_public_http_url, + validate_public_http_redirect_url, +) # Content-Type to file extension mapping CONTENT_TYPE_TO_EXTENSION: dict[str, str] = { @@ -33,6 +39,8 @@ "image/svg+xml": ".svg", } +REDIRECT_STATUS_CODES: set[int] = {301, 302, 303, 307, 308} + def _extension_from_path(url: str) -> str | None: """Extract a recognised file extension from the URL path.""" @@ -63,9 +71,9 @@ async def resolve_file_extension_async(url: str) -> str | None: 2. If that fails, send a HEAD request and read Content-Type. 3. Return None if neither method produces a supported extension. """ - validate_public_http_url(url, field="source_url") + safe_url = get_safe_public_http_url(url, field="source_url") - ext = _extension_from_path(url) + ext = _extension_from_path(safe_url) if ext: return ext @@ -73,17 +81,36 @@ async def resolve_file_extension_async(url: str) -> str | None: from shared.utils.http_clients import get_async_client client = get_async_client() - response = await client.head(url, follow_redirects=True) + response = None + request_url: SafePublicHTTPURL = safe_url + for _ in range(MAX_SAFE_REDIRECTS + 1): + response = await client.head(request_url, follow_redirects=False) + if response.status_code not in REDIRECT_STATUS_CODES: + break + + location = response.headers.get("location") + if not location: + break + request_url = validate_public_http_redirect_url( + request_url, + location, + field="source_url", + ) + + if response is None: + return None content_type = response.headers.get("content-type") ext = _extension_from_content_type(content_type) if ext: logger.info( - f"Resolved file extension from Content-Type header: {ext} (url={url})" + f"Resolved file extension from Content-Type header: {ext}" ) return ext + except ValidationException: + raise except Exception as exc: logger.warning( - f"HEAD request failed for URL file type detection: {exc} (url={url})" + f"HEAD request failed for URL file type detection: {exc}" ) return None @@ -95,9 +122,9 @@ def resolve_file_extension_sync(url: str) -> str | None: Same logic as async variant but uses the shared sync httpx client. """ - validate_public_http_url(url, field="source_url") + safe_url = get_safe_public_http_url(url, field="source_url") - ext = _extension_from_path(url) + ext = _extension_from_path(safe_url) if ext: return ext @@ -105,17 +132,36 @@ def resolve_file_extension_sync(url: str) -> str | None: from shared.utils.http_clients import get_sync_client client = get_sync_client() - response = client.head(url, follow_redirects=True) + response = None + request_url: SafePublicHTTPURL = safe_url + for _ in range(MAX_SAFE_REDIRECTS + 1): + response = client.head(request_url, follow_redirects=False) + if response.status_code not in REDIRECT_STATUS_CODES: + break + + location = response.headers.get("location") + if not location: + break + request_url = validate_public_http_redirect_url( + request_url, + location, + field="source_url", + ) + + if response is None: + return None content_type = response.headers.get("content-type") ext = _extension_from_content_type(content_type) if ext: logger.info( - f"Resolved file extension from Content-Type header: {ext} (url={url})" + f"Resolved file extension from Content-Type header: {ext}" ) return ext + except ValidationException: + raise except Exception as exc: logger.warning( - f"HEAD request failed for URL file type detection: {exc} (url={url})" + f"HEAD request failed for URL file type detection: {exc}" ) return None diff --git a/packages/shared-python/shared/utils/url_security.py b/packages/shared-python/shared/utils/url_security.py index b96da5398..34154417b 100644 --- a/packages/shared-python/shared/utils/url_security.py +++ b/packages/shared-python/shared/utils/url_security.py @@ -2,12 +2,13 @@ import ipaddress import socket from typing import cast -from urllib.parse import urlparse +from urllib.parse import urljoin, urlparse from shared.core.exceptions.domain_exceptions import ValidationException AddressInfo = tuple[int, int, int, str, tuple[str, ...]] AllowedIPAddress = ipaddress.IPv4Address | ipaddress.IPv6Address +MAX_SAFE_REDIRECTS = 5 class URLSecurityError(ValueError): @@ -26,18 +27,26 @@ class InvalidResolvedAddressError(URLSecurityError): """Raised when DNS returns an invalid IP address.""" -def validate_public_http_url(url: str, field: str = "url") -> None: - """Reject URL inputs that could target internal networks or local services.""" - parsed_url = urlparse(url) - if parsed_url.scheme not in {"http", "https"}: - raise _build_url_validation_error(field, "URL must use http or https") +class UnsupportedURLSchemeError(URLSecurityError): + """Raised when a URL uses a disallowed scheme.""" - hostname = parsed_url.hostname - if not hostname: - raise _build_url_validation_error(field, "URL must include a hostname") +class MissingURLHostnameError(URLSecurityError): + """Raised when a URL does not include a hostname.""" + + +class SafePublicHTTPURL(str): + """A URL string that has passed public HTTP SSRF validation.""" + + +def validate_public_http_url(url: str, field: str = "url") -> None: + """Reject URL inputs that could target internal networks or local services.""" try: - resolve_public_hostname(hostname) + _validate_public_http_url(url) + except UnsupportedURLSchemeError as exc: + raise _build_url_validation_error(field, "URL must use http or https") from exc + except MissingURLHostnameError as exc: + raise _build_url_validation_error(field, "URL must include a hostname") from exc except HostnameResolutionError as exc: raise _build_url_validation_error(field, "URL hostname could not be resolved") from exc except InvalidResolvedAddressError as exc: @@ -46,6 +55,30 @@ def validate_public_http_url(url: str, field: str = "url") -> None: raise _build_url_validation_error(field, "URL host is not allowed") from exc +def validate_public_http_redirect_url( + url: str, + redirect_url: str, + field: str = "url", +) -> SafePublicHTTPURL: + """Resolve and validate an HTTP redirect target before following it.""" + resolved_url = urljoin(url, redirect_url) + validate_public_http_url(resolved_url, field=field) + return SafePublicHTTPURL(resolved_url) + + +def get_safe_public_http_url(url: str, field: str = "url") -> SafePublicHTTPURL: + """Return a validated public HTTP URL for outbound HTTP clients.""" + validate_public_http_url(url, field=field) + return SafePublicHTTPURL(url) + + +def get_safe_redirect_url(url: str, redirect_url: str) -> str: + """Resolve and validate an HTTP redirect target for internal network callers.""" + resolved_url = urljoin(url, redirect_url) + _validate_public_http_url(resolved_url) + return resolved_url + + def resolve_public_hostname(hostname: str) -> str: """Resolve a hostname to a public IP address and return the pinned IP.""" _ensure_hostname_is_allowed(hostname) @@ -59,6 +92,18 @@ def resolve_public_hostname(hostname: str) -> str: return _select_public_ip_address(hostname, address_infos) +def _validate_public_http_url(url: str) -> None: + parsed_url = urlparse(url) + if parsed_url.scheme not in {"http", "https"}: + raise UnsupportedURLSchemeError(f"Unsupported URL scheme: {parsed_url.scheme}") + + hostname = parsed_url.hostname + if not hostname: + raise MissingURLHostnameError("URL must include a hostname") + + resolve_public_hostname(hostname) + + async def resolve_public_hostname_async(hostname: str) -> str: """Resolve a hostname to a public IP address asynchronously.""" _ensure_hostname_is_allowed(hostname) From a2c36d84d06c2e64f0ea7f29d4e351a74b9d5ec1 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Mon, 4 May 2026 10:19:57 +0000 Subject: [PATCH 02/32] refactor: remove unused api key hash version --- .../alembic/versions/f6a7b8c9d0e1_key_api_key_hashes.py | 8 +------- apps/api/tests/support/contract_database.py | 3 --- packages/shared-python/shared/models/database/api_key.py | 3 --- 3 files changed, 1 insertion(+), 13 deletions(-) diff --git a/apps/api/alembic/versions/f6a7b8c9d0e1_key_api_key_hashes.py b/apps/api/alembic/versions/f6a7b8c9d0e1_key_api_key_hashes.py index cb7688d90..cd63a6734 100644 --- a/apps/api/alembic/versions/f6a7b8c9d0e1_key_api_key_hashes.py +++ b/apps/api/alembic/versions/f6a7b8c9d0e1_key_api_key_hashes.py @@ -8,7 +8,6 @@ from typing import Sequence, Union -import sqlalchemy as sa from alembic import op @@ -20,12 +19,7 @@ def upgrade() -> None: op.execute("UPDATE api_keys SET is_active = false") - op.add_column( - "api_keys", - sa.Column("hash_version", sa.String(length=16), nullable=False, server_default="hmac-v1"), - ) - op.alter_column("api_keys", "hash_version", server_default=None) def downgrade() -> None: - op.drop_column("api_keys", "hash_version") + pass diff --git a/apps/api/tests/support/contract_database.py b/apps/api/tests/support/contract_database.py index f704ff41c..28f955e10 100644 --- a/apps/api/tests/support/contract_database.py +++ b/apps/api/tests/support/contract_database.py @@ -148,7 +148,6 @@ async def insert_authenticated_user( id, user_id, key_hash, - hash_version, key_mask, name, enabled_modules, @@ -158,7 +157,6 @@ async def insert_authenticated_user( :id, :user_id, :key_hash, - :hash_version, :key_mask, :name, CAST(:enabled_modules AS JSON), @@ -170,7 +168,6 @@ async def insert_authenticated_user( "id": api_key_id, "user_id": user_id, "key_hash": api_key_hash, - "hash_version": "hmac-v1", "key_mask": f"{api_key[:8]}...{api_key[-4:]}", "name": f"Contract API Key {user_id}", "enabled_modules": json.dumps(enabled_modules or ["all"]), diff --git a/packages/shared-python/shared/models/database/api_key.py b/packages/shared-python/shared/models/database/api_key.py index 55028c1f6..a70b24f9f 100644 --- a/packages/shared-python/shared/models/database/api_key.py +++ b/packages/shared-python/shared/models/database/api_key.py @@ -29,9 +29,6 @@ class APIKey(Base): key_hash: Mapped[str] = mapped_column( String(255), nullable=False, index=True ) # Encrypted storage - hash_version: Mapped[str] = mapped_column( - String(16), default="hmac-v1", nullable=False - ) key_mask: Mapped[str] = mapped_column( String(50), nullable=False ) # Masked API Key (for display) From b77b01db3f5acbf909fad407cf5d085e13356537 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Tue, 5 May 2026 09:41:55 +0000 Subject: [PATCH 03/32] refactor: remove database health endpoints --- apps/api/app/api/v1/health.py | 14 -- .../contract/test_database_health_contract.py | 55 -------- .../shared-python/shared/core/database.py | 128 ------------------ 3 files changed, 197 deletions(-) diff --git a/apps/api/app/api/v1/health.py b/apps/api/app/api/v1/health.py index f01472202..43837721f 100644 --- a/apps/api/app/api/v1/health.py +++ b/apps/api/app/api/v1/health.py @@ -5,7 +5,6 @@ from fastapi import APIRouter from shared.core.database import ( - get_database_health, get_database_performance, prewarm_connection_pool, ) @@ -13,19 +12,6 @@ router = APIRouter() -@router.get("/database/health") -async def check_database_health(): - """Check database health status.""" - health = await get_database_health() - if "error" in health: - return { - "status": health.get("status", "unhealthy"), - "error": "Database health check failed", - "last_check": health.get("last_check"), - } - return health - - @router.get("/database/performance") async def get_database_performance_stats(): """Return database performance statistics.""" diff --git a/apps/api/tests/contract/test_database_health_contract.py b/apps/api/tests/contract/test_database_health_contract.py index 2456ec298..135bb0833 100644 --- a/apps/api/tests/contract/test_database_health_contract.py +++ b/apps/api/tests/contract/test_database_health_contract.py @@ -4,61 +4,6 @@ import pytest from httpx import AsyncClient -from pytest import MonkeyPatch - - -@pytest.mark.asyncio -async def test_should_return_the_database_health_payload_shape( - api_client_factory: Callable[[], AbstractAsyncContextManager[AsyncClient]], -) -> None: - async with api_client_factory() as api_client: - response = await api_client.get("/api/v1/health/database/health") - - assert response.status_code == 200 - - response_json = cast(dict[str, object], response.json()) - pool_status = cast(dict[str, object], response_json["pool_status"]) - - assert response_json["status"] == "healthy" - assert isinstance(response_json["response_time_ms"], float | int) - assert response_json["last_check"] - assert set(pool_status) == { - "size", - "checked_in", - "checked_out", - "overflow", - "invalid", - } - - -@pytest.mark.asyncio -async def test_should_sanitize_database_health_errors_before_returning_them( - api_client_factory: Callable[[], AbstractAsyncContextManager[AsyncClient]], - monkeypatch: MonkeyPatch, -) -> None: - async def _get_database_health() -> dict[str, object]: - return { - "status": "unhealthy", - "error": "Database health check failed", - "last_check": None, - } - - async with api_client_factory() as api_client: - from app.api.v1 import health as health_route_module - - monkeypatch.setattr( - health_route_module, - "get_database_health", - _get_database_health, - ) - response = await api_client.get("/api/v1/health/database/health") - - assert response.status_code == 200 - assert response.json() == { - "status": "unhealthy", - "error": "Database health check failed", - "last_check": None, - } @pytest.mark.asyncio diff --git a/packages/shared-python/shared/core/database.py b/packages/shared-python/shared/core/database.py index acdf08654..8087d79f2 100644 --- a/packages/shared-python/shared/core/database.py +++ b/packages/shared-python/shared/core/database.py @@ -1,7 +1,6 @@ import asyncio import logging import os -import time from contextlib import asynccontextmanager from datetime import datetime from typing import Any, AsyncGenerator, Awaitable, Callable, TypeVar @@ -144,133 +143,6 @@ async def create_tables(): await conn.run_sync(Base.metadata.create_all) -# Connection-pool monitoring and health checks. -class DatabaseHealthChecker: - """Database health checker.""" - - def __init__(self, engine: AsyncEngine): - self.engine = engine - self.last_check: datetime | None = None - self.is_healthy = False - - def _pool_metric(self, name: str) -> int: - metric = getattr(self.engine.pool, name, None) - if callable(metric): - value = metric() - return int(value) if isinstance(value, (int, float)) else 0 - return 0 - - async def check_health(self) -> dict[str, object]: - """Check database connection health.""" - try: - start_time = time.time() - async with self.engine.begin() as conn: - # Run a simple query to validate connectivity. - result = await conn.execute(text("SELECT 1 as health_check")) - row = result.fetchone() - - if row and row[0] == 1: - self.is_healthy = True - self.last_check = datetime.now() - - # Collect current connection-pool status. - pool_status = self.get_pool_status() - - return { - "status": "healthy", - "response_time_ms": round((time.time() - start_time) * 1000, 2), - "last_check": self.last_check.isoformat(), - "pool_status": pool_status, - } - else: - self.is_healthy = False - return { - "status": "unhealthy", - "error": "Health check query failed", - "last_check": ( - self.last_check.isoformat() if self.last_check else None - ), - } - except Exception as e: - self.is_healthy = False - logger.error(f"Database health check failed: {e}") - return { - "status": "unhealthy", - "error": "Database health check failed", - "last_check": self.last_check.isoformat() if self.last_check else None, - } - - def get_pool_status(self) -> dict[str, int]: - """Return connection-pool status details.""" - status = { - "size": self._pool_metric("size"), - "checked_in": self._pool_metric("checkedin"), - "checked_out": self._pool_metric("checkedout"), - "overflow": self._pool_metric("overflow"), - } - status["invalid"] = self._pool_metric("invalid") - return status - - async def get_database_info(self) -> dict[str, object]: - """Return database metadata and connection status.""" - try: - async with self.engine.begin() as conn: - # Read the database version. - version_result = await conn.execute(text("SELECT version()")) - version_row = version_result.fetchone() - if version_row is None: - return {"error": "Failed to read database version"} - version = version_row[0] - - # Read the current active-connection count. - connections_result = await conn.execute( - text(""" - SELECT count(*) as active_connections - FROM pg_stat_activity - WHERE state = 'active' - """) - ) - active_connections_row = connections_result.fetchone() - if active_connections_row is None: - return {"error": "Failed to read active connection count"} - active_connections = active_connections_row[0] - - # Read the current database size. - size_result = await conn.execute( - text(""" - SELECT pg_size_pretty(pg_database_size(current_database())) as db_size - """) - ) - db_size_row = size_result.fetchone() - if db_size_row is None: - return {"error": "Failed to read database size"} - db_size = db_size_row[0] - - return { - "version": version, - "active_connections": active_connections, - "database_size": db_size, - "pool_status": self.get_pool_status(), - } - except Exception as e: - logger.error(f"Failed to get database info: {e}") - return {"error": "Database information unavailable"} - - -# Shared health-checker instance. -db_health_checker = DatabaseHealthChecker(engine) - - -async def get_database_health() -> dict: - """Return database health status.""" - return await db_health_checker.check_health() - - -async def get_database_info() -> dict: - """Return database information.""" - return await db_health_checker.get_database_info() - - # Database retry helpers. class DatabaseRetryManager: """Database retry manager.""" From be959c5dfa4ef65efe734cbf1bc54d234d095696 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Tue, 5 May 2026 09:59:21 +0000 Subject: [PATCH 04/32] refactor: extract pinned outbound url helpers --- apps/api/app/api/v1/routes/jobs.py | 6 +- apps/api/app/api/v1/routes/s3_events.py | 83 ++++---------- .../tests/contract/test_s3_event_contract.py | 6 +- .../test_webhook_recovery_contract.py | 2 +- .../shared/services/webhook/dispatcher.py | 102 ++++++------------ ...validator.py => outbound_url_validator.py} | 43 ++++---- .../services/webhook/pinned_outbound_http.py | 87 +++++++++++++++ .../services/webhook/qstash_publisher.py | 4 +- 8 files changed, 176 insertions(+), 157 deletions(-) rename packages/shared-python/shared/services/webhook/{validator.py => outbound_url_validator.py} (63%) create mode 100644 packages/shared-python/shared/services/webhook/pinned_outbound_http.py diff --git a/apps/api/app/api/v1/routes/jobs.py b/apps/api/app/api/v1/routes/jobs.py index c4e9bcae6..e9e0b77b8 100644 --- a/apps/api/app/api/v1/routes/jobs.py +++ b/apps/api/app/api/v1/routes/jobs.py @@ -53,7 +53,9 @@ StandardErrorObject, ) from shared.services.storage.file_upload_service import FileUploadService -from shared.services.webhook.validator import validate_webhook_url_async +from shared.services.webhook.outbound_url_validator import ( + validate_outbound_url_async, +) from shared.utils.error_details import normalize_error_details from shared.utils.url_file_type import resolve_file_extension_async @@ -320,7 +322,7 @@ async def create_job( # pyright: ignore[reportGeneralTypeIssues] if payload.webhook: # Check for URL validity if payload.webhook.url: - validation_result = await validate_webhook_url_async( + validation_result = await validate_outbound_url_async( payload.webhook.url ) if not validation_result.is_valid: diff --git a/apps/api/app/api/v1/routes/s3_events.py b/apps/api/app/api/v1/routes/s3_events.py index 9b0897ca0..61e1abab8 100644 --- a/apps/api/app/api/v1/routes/s3_events.py +++ b/apps/api/app/api/v1/routes/s3_events.py @@ -5,11 +5,8 @@ import base64 import json import os -import socket from typing import Any, Dict -import aiohttp -from aiohttp.abc import AbstractResolver from app.repositories.job_repository import JobRepository from app.services.knowledge.kb_orchestrator import KBOrchestrator from app.services.state_machine import JobStateMachine @@ -21,8 +18,10 @@ from shared.core.state_machine.states import JobStatus from shared.models.schemas.oss_event import OSSEvent from shared.models.schemas.s3_event import S3Event -from shared.services.webhook.validator import validate_webhook_url_async -from shared.utils.url_security import SafePublicHTTPURL +from shared.services.webhook.pinned_outbound_http import ( + send_pinned_outbound_request, +) +from shared.services.webhook.outbound_url_validator import validate_outbound_url_async router = APIRouter(tags=["Internal"]) @@ -267,7 +266,7 @@ async def handle_sns_event(body: bytes): async def confirm_sns_subscription(subscribe_url: str) -> dict[str, str]: """Confirm an SNS subscription after SSRF validation and IP pinning.""" - validation = await validate_webhook_url_async(subscribe_url) + validation = await validate_outbound_url_async(subscribe_url) if not validation.is_valid: logger.warning( f"SNS subscription confirmation URL failed validation: {validation.error_message}" @@ -279,66 +278,30 @@ async def confirm_sns_subscription(subscribe_url: str) -> dict[str, str]: return {"message": "SNS subscription confirmation failed"} try: - validated_subscribe_url = SafePublicHTTPURL(subscribe_url) - connector = aiohttp.TCPConnector( - resolver=_PinnedSNSResolver(validation.validated_ip), + response = await send_pinned_outbound_request( + method="GET", + url=subscribe_url, + pinned_ip=validation.validated_ip, + timeout_seconds=SNS_SUBSCRIPTION_TIMEOUT_SECONDS, ) - timeout = aiohttp.ClientTimeout(total=SNS_SUBSCRIPTION_TIMEOUT_SECONDS) - async with aiohttp.ClientSession( - connector=connector, - timeout=timeout, - ) as session: - async with session.get( - validated_subscribe_url, - allow_redirects=False, - ) as response: - if response.status == 200: - logger.info("SNS subscription confirmed successfully") - return {"message": "SNS subscription confirmed"} - - if 300 <= response.status < 400: - logger.warning( - f"SNS subscription confirmation redirect blocked, status={response.status}" - ) - else: - logger.error( - f"SNS subscription confirmation failed, status={response.status}" - ) - return {"message": "SNS subscription confirmation failed"} + if response.status == 200: + logger.info("SNS subscription confirmed successfully") + return {"message": "SNS subscription confirmed"} + + if 300 <= response.status < 400: + logger.warning( + f"SNS subscription confirmation redirect blocked, status={response.status}" + ) + else: + logger.error( + f"SNS subscription confirmation failed, status={response.status}" + ) + return {"message": "SNS subscription confirmation failed"} except Exception as e: logger.error(f"Failed to reach the SNS confirmation URL: {e}") return {"message": "SNS subscription confirmation failed"} -class _PinnedSNSResolver(AbstractResolver): - """Resolver that pins SNS confirmation to a pre-validated public IP.""" - - def __init__(self, pinned_ip: str) -> None: - self.pinned_ip = pinned_ip - - async def resolve( - self, - host: str, - port: int = 0, - family: int = socket.AF_INET, - ) -> list[dict[str, Any]]: - parsed_ip = self.pinned_ip - pinned_family = socket.AF_INET6 if ":" in parsed_ip else socket.AF_INET - return [ - { - "hostname": host, - "host": parsed_ip, - "port": port, - "family": pinned_family, - "proto": 0, - "flags": socket.AI_NUMERICHOST, - } - ] - - async def close(self) -> None: - pass - - async def handle_minio_event(body: bytes, auth_token: str): """ Handle a MinIO webhook event. diff --git a/apps/api/tests/contract/test_s3_event_contract.py b/apps/api/tests/contract/test_s3_event_contract.py index 5b07b2fb7..7464307b5 100644 --- a/apps/api/tests/contract/test_s3_event_contract.py +++ b/apps/api/tests/contract/test_s3_event_contract.py @@ -196,10 +196,12 @@ def get(self, url: str, *args: object, **kwargs: object) -> object: raise AssertionError("private SNS confirmation URL should not be requested") async with api_client_factory() as api_client: - s3_events_module = importlib.import_module("app.api.v1.routes.s3_events") monkeypatch.setattr(socket, "getaddrinfo", resolve_private_address) + pinned_http_module = importlib.import_module( + "shared.services.webhook.pinned_outbound_http" + ) monkeypatch.setattr( - s3_events_module.aiohttp, + pinned_http_module.aiohttp, "ClientSession", _UnexpectedSession, ) diff --git a/apps/worker/tests/contract/test_webhook_recovery_contract.py b/apps/worker/tests/contract/test_webhook_recovery_contract.py index ecb469c58..3afab298f 100644 --- a/apps/worker/tests/contract/test_webhook_recovery_contract.py +++ b/apps/worker/tests/contract/test_webhook_recovery_contract.py @@ -119,7 +119,7 @@ def publish(self, **kwargs: Any) -> SimpleNamespace: ) monkeypatch.setattr( qstash_publisher, - "validate_webhook_url", + "validate_outbound_url", lambda url: SimpleNamespace( is_valid=True, error_message=None, diff --git a/packages/shared-python/shared/services/webhook/dispatcher.py b/packages/shared-python/shared/services/webhook/dispatcher.py index bb0eb71ca..8e2e4171c 100644 --- a/packages/shared-python/shared/services/webhook/dispatcher.py +++ b/packages/shared-python/shared/services/webhook/dispatcher.py @@ -9,15 +9,12 @@ import hashlib import hmac import json -import socket import threading import time import uuid from datetime import datetime, timezone from typing import Any, Dict, Optional, Tuple -import aiohttp -from aiohttp.abc import AbstractResolver from loguru import logger from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession @@ -32,9 +29,12 @@ from shared.models.database.job import Job from shared.models.database.webhook import WebhookEvent, WebhookEventStatus from shared.models.database.webhook_log import WebhookLog -from shared.services.webhook.validator import ( - WebhookValidationResult, - validate_webhook_url_async, +from shared.services.webhook.pinned_outbound_http import ( + send_pinned_outbound_request, +) +from shared.services.webhook.outbound_url_validator import ( + OutboundURLValidationResult, + validate_outbound_url_async, ) # Constants @@ -166,7 +166,7 @@ async def _send_webhook( attempt_id = str(uuid.uuid4()) # SSRF Protection - validation: WebhookValidationResult = await validate_webhook_url_async( + validation: OutboundURLValidationResult = await validate_outbound_url_async( event.target_url ) if not validation.is_valid: @@ -226,71 +226,39 @@ async def _get_job_owner(job_id: str) -> Optional[str]: success = False try: - # IP Pinning: Use a custom resolver that returns ONLY the pre-validated IP. - # This eliminates the DNS rebinding TOCTOU window — aiohttp will connect - # to the pinned IP while the Host header preserves the original hostname. pinned_ip = validation.validated_ip if not pinned_ip: return False, 400, 0, "SSRF validation did not return a pinned IP" - # Detect address family from the pinned IP - pinned_family: int = socket.AF_INET6 if ":" in pinned_ip else socket.AF_INET - - class PinnedResolver(AbstractResolver): - """Resolver that always returns the pre-validated IP address.""" - - async def resolve( - self, host: str, port: int = 0, family: int = socket.AF_INET - ) -> list[dict[str, Any]]: - return [ - { - "hostname": host, - "host": pinned_ip, - "port": port, - "family": pinned_family, - "proto": 0, - "flags": socket.AI_NUMERICHOST, - } - ] - - async def close(self) -> None: - pass - - connector = aiohttp.TCPConnector( - resolver=PinnedResolver(), - # Disable redirect following to prevent redirect-based SSRF - # (attacker returns 302 → http://169.254.169.254/...) + response = await send_pinned_outbound_request( + method="POST", + url=event.target_url, + pinned_ip=pinned_ip, + timeout_seconds=HTTP_TIMEOUT_SECONDS, + headers=headers, + json_body=enriched_payload, ) - async with aiohttp.ClientSession(connector=connector) as session: - async with session.post( - event.target_url, - json=enriched_payload, - headers=headers, - timeout=aiohttp.ClientTimeout(total=HTTP_TIMEOUT_SECONDS), - allow_redirects=False, # Block redirect-based SSRF - ) as response: - duration_ms = int((time.time() - start_time) * 1000) - status_code = response.status - - # Treat 3xx as non-success (redirect-based SSRF prevention) - if 200 <= response.status < 300: - logger.info( - f"Webhook delivered: event_id={event.id}, status={response.status}" - ) - success = True - elif 300 <= response.status < 400: - logger.warning( - f"Webhook redirect blocked (SSRF protection): " - f"event_id={event.id}, status={response.status}" - ) - error_message = f"Redirect blocked: HTTP {response.status}" - success = False - else: - logger.warning( - f"Webhook failed: event_id={event.id}, status={response.status}" - ) - error_message = f"HTTP {response.status}" - success = False + duration_ms = int((time.time() - start_time) * 1000) + status_code = response.status + + if 200 <= response.status < 300: + logger.info( + f"Webhook delivered: event_id={event.id}, status={response.status}" + ) + success = True + elif 300 <= response.status < 400: + logger.warning( + f"Webhook redirect blocked (SSRF protection): " + f"event_id={event.id}, status={response.status}" + ) + error_message = f"Redirect blocked: HTTP {response.status}" + success = False + else: + logger.warning( + f"Webhook failed: event_id={event.id}, status={response.status}" + ) + error_message = f"HTTP {response.status}" + success = False except asyncio.TimeoutError: duration_ms = int((time.time() - start_time) * 1000) diff --git a/packages/shared-python/shared/services/webhook/validator.py b/packages/shared-python/shared/services/webhook/outbound_url_validator.py similarity index 63% rename from packages/shared-python/shared/services/webhook/validator.py rename to packages/shared-python/shared/services/webhook/outbound_url_validator.py index 23f6c949f..40a203b86 100644 --- a/packages/shared-python/shared/services/webhook/validator.py +++ b/packages/shared-python/shared/services/webhook/outbound_url_validator.py @@ -1,7 +1,7 @@ """ -Webhook Validation Utilities +Outbound URL Validation Utilities -SSRF protection via shared DNS/IP validation + IP pinning. +Shared SSRF protection for outbound HTTP targets via DNS/IP validation + IP pinning. """ from dataclasses import dataclass @@ -12,12 +12,9 @@ from shared.utils.url_security import resolve_public_hostname, resolve_public_hostname_async -# ── Integration wrapper for our dispatcher ──────────────────────────── - - @dataclass -class WebhookValidationResult: - """Result of webhook URL validation, including pinned IP for anti-DNS-rebinding.""" +class OutboundURLValidationResult: + """Result of outbound URL validation, including a pinned IP address.""" is_valid: bool error_message: Optional[str] = None @@ -25,70 +22,70 @@ class WebhookValidationResult: hostname: Optional[str] = None -async def validate_webhook_url_async(url: str) -> WebhookValidationResult: +async def validate_outbound_url_async(url: str) -> OutboundURLValidationResult: """ - Async webhook URL validation with SSRF protection and IP pinning. + Async outbound URL validation with SSRF protection and IP pinning. - Returns WebhookValidationResult with validated_ip for the dispatcher - to pin the connection to, eliminating the DNS rebinding TOCTOU window. + Returns OutboundURLValidationResult with a pinned IP address, eliminating + the DNS rebinding TOCTOU window for later outbound requests. """ try: parsed = urlparse(url) is_dev: bool = app_config.ENVIRONMENT.lower() in ("dev", "development", "local") allowed_schemes: list[str] = ["https"] if not is_dev else ["https", "http"] if parsed.scheme not in allowed_schemes: - return WebhookValidationResult( + return OutboundURLValidationResult( is_valid=False, error_message=f"Invalid scheme: {parsed.scheme}. Must be HTTPS.", ) hostname: Optional[str] = parsed.hostname if not hostname: - return WebhookValidationResult( + return OutboundURLValidationResult( is_valid=False, error_message="URL must have a hostname" ) validated_ip: str = await resolve_public_hostname_async(hostname) - return WebhookValidationResult( + return OutboundURLValidationResult( is_valid=True, validated_ip=validated_ip, hostname=hostname, ) except ValueError as exc: - return WebhookValidationResult(is_valid=False, error_message=str(exc)) + return OutboundURLValidationResult(is_valid=False, error_message=str(exc)) except Exception as exc: - return WebhookValidationResult( + return OutboundURLValidationResult( is_valid=False, error_message=f"URL validation failed: {exc}", ) -def validate_webhook_url(url: str) -> WebhookValidationResult: - """Sync webhook URL validation with SSRF checks.""" +def validate_outbound_url(url: str) -> OutboundURLValidationResult: + """Sync outbound URL validation with SSRF checks.""" try: parsed = urlparse(url) is_dev: bool = app_config.ENVIRONMENT.lower() in ("dev", "development", "local") allowed_schemes: list[str] = ["https"] if not is_dev else ["https", "http"] if parsed.scheme not in allowed_schemes: - return WebhookValidationResult( + return OutboundURLValidationResult( is_valid=False, error_message=f"Invalid scheme: {parsed.scheme}. Must be HTTPS.", ) hostname: Optional[str] = parsed.hostname if not hostname: - return WebhookValidationResult( + return OutboundURLValidationResult( is_valid=False, error_message="URL must have a hostname" ) validated_ip: str = resolve_public_hostname(hostname) - return WebhookValidationResult( + return OutboundURLValidationResult( is_valid=True, validated_ip=validated_ip, hostname=hostname, ) except ValueError as exc: - return WebhookValidationResult(is_valid=False, error_message=str(exc)) + return OutboundURLValidationResult(is_valid=False, error_message=str(exc)) except Exception as exc: - return WebhookValidationResult( + return OutboundURLValidationResult( is_valid=False, error_message=f"URL validation failed: {exc}", ) diff --git a/packages/shared-python/shared/services/webhook/pinned_outbound_http.py b/packages/shared-python/shared/services/webhook/pinned_outbound_http.py new file mode 100644 index 000000000..6b4db161a --- /dev/null +++ b/packages/shared-python/shared/services/webhook/pinned_outbound_http.py @@ -0,0 +1,87 @@ +""" +Pinned outbound HTTP helpers. + +Shared infrastructure for outbound requests that must connect to a +pre-validated public IP address and block redirect-based SSRF. +""" + +from __future__ import annotations + +import socket +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Any + +import aiohttp +from aiohttp.abc import AbstractResolver + +from shared.utils.url_security import SafePublicHTTPURL + + +@dataclass(frozen=True) +class PinnedOutboundResponse: + """Minimal response metadata for pinned outbound HTTP requests.""" + + status: int + + +class PinnedIPResolver(AbstractResolver): + """Resolver that always returns the supplied pinned IP address.""" + + def __init__(self, pinned_ip: str) -> None: + self.pinned_ip = pinned_ip + + async def resolve( + self, + host: str, + port: int = 0, + family: int = socket.AF_INET, + ) -> list[dict[str, Any]]: + pinned_family: int = ( + socket.AF_INET6 if ":" in self.pinned_ip else socket.AF_INET + ) + return [ + { + "hostname": host, + "host": self.pinned_ip, + "port": port, + "family": pinned_family, + "proto": 0, + "flags": socket.AI_NUMERICHOST, + } + ] + + async def close(self) -> None: + pass + + +async def send_pinned_outbound_request( + *, + method: str, + url: str, + pinned_ip: str, + timeout_seconds: float, + headers: Mapping[str, str] | None = None, + json_body: Any | None = None, +) -> PinnedOutboundResponse: + """ + Send an outbound HTTP request through a resolver pinned to a validated IP. + + Redirects are always blocked to prevent redirect-based SSRF. + """ + validated_url = SafePublicHTTPURL(url) + connector = aiohttp.TCPConnector(resolver=PinnedIPResolver(pinned_ip)) + timeout = aiohttp.ClientTimeout(total=timeout_seconds) + + async with aiohttp.ClientSession( + connector=connector, + timeout=timeout, + ) as session: + async with session.request( + method=method, + url=validated_url, + headers=headers, + json=json_body, + allow_redirects=False, + ) as response: + return PinnedOutboundResponse(status=response.status) diff --git a/packages/shared-python/shared/services/webhook/qstash_publisher.py b/packages/shared-python/shared/services/webhook/qstash_publisher.py index da8ac61a6..6bb4d137d 100644 --- a/packages/shared-python/shared/services/webhook/qstash_publisher.py +++ b/packages/shared-python/shared/services/webhook/qstash_publisher.py @@ -22,7 +22,7 @@ from shared.core.config import app_config from shared.core.exceptions.domain_exceptions import QStashServiceException from shared.models.database.webhook import WebhookEventStatus -from shared.services.webhook.validator import validate_webhook_url +from shared.services.webhook.outbound_url_validator import validate_outbound_url class QStashWebhookPublisher: @@ -84,7 +84,7 @@ def publish_event(self, event_id: str) -> Optional[str]: return None # SSRF pre-validation - validation = validate_webhook_url(event.target_url) + validation = validate_outbound_url(event.target_url) if not validation.is_valid: logger.warning( f"QStash publish: SSRF validation failed for event {event_id}: " From 4e5979f97e4ab1985385f617646ab071d33657f2 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Tue, 5 May 2026 10:07:02 +0000 Subject: [PATCH 05/32] refactor: move outbound helpers to shared utils --- apps/api/app/api/v1/routes/jobs.py | 2 +- apps/api/app/api/v1/routes/s3_events.py | 4 ++-- apps/api/tests/contract/test_s3_event_contract.py | 2 +- packages/shared-python/shared/services/webhook/dispatcher.py | 4 ++-- .../shared-python/shared/services/webhook/qstash_publisher.py | 2 +- .../{services/webhook => utils}/outbound_url_validator.py | 0 .../{services/webhook => utils}/pinned_outbound_http.py | 0 7 files changed, 7 insertions(+), 7 deletions(-) rename packages/shared-python/shared/{services/webhook => utils}/outbound_url_validator.py (100%) rename packages/shared-python/shared/{services/webhook => utils}/pinned_outbound_http.py (100%) diff --git a/apps/api/app/api/v1/routes/jobs.py b/apps/api/app/api/v1/routes/jobs.py index e9e0b77b8..40f0dee33 100644 --- a/apps/api/app/api/v1/routes/jobs.py +++ b/apps/api/app/api/v1/routes/jobs.py @@ -53,7 +53,7 @@ StandardErrorObject, ) from shared.services.storage.file_upload_service import FileUploadService -from shared.services.webhook.outbound_url_validator import ( +from shared.utils.outbound_url_validator import ( validate_outbound_url_async, ) from shared.utils.error_details import normalize_error_details diff --git a/apps/api/app/api/v1/routes/s3_events.py b/apps/api/app/api/v1/routes/s3_events.py index 61e1abab8..40d0bb99e 100644 --- a/apps/api/app/api/v1/routes/s3_events.py +++ b/apps/api/app/api/v1/routes/s3_events.py @@ -18,10 +18,10 @@ from shared.core.state_machine.states import JobStatus from shared.models.schemas.oss_event import OSSEvent from shared.models.schemas.s3_event import S3Event -from shared.services.webhook.pinned_outbound_http import ( +from shared.utils.pinned_outbound_http import ( send_pinned_outbound_request, ) -from shared.services.webhook.outbound_url_validator import validate_outbound_url_async +from shared.utils.outbound_url_validator import validate_outbound_url_async router = APIRouter(tags=["Internal"]) diff --git a/apps/api/tests/contract/test_s3_event_contract.py b/apps/api/tests/contract/test_s3_event_contract.py index 7464307b5..949722136 100644 --- a/apps/api/tests/contract/test_s3_event_contract.py +++ b/apps/api/tests/contract/test_s3_event_contract.py @@ -198,7 +198,7 @@ def get(self, url: str, *args: object, **kwargs: object) -> object: async with api_client_factory() as api_client: monkeypatch.setattr(socket, "getaddrinfo", resolve_private_address) pinned_http_module = importlib.import_module( - "shared.services.webhook.pinned_outbound_http" + "shared.utils.pinned_outbound_http" ) monkeypatch.setattr( pinned_http_module.aiohttp, diff --git a/packages/shared-python/shared/services/webhook/dispatcher.py b/packages/shared-python/shared/services/webhook/dispatcher.py index 8e2e4171c..c3d166e58 100644 --- a/packages/shared-python/shared/services/webhook/dispatcher.py +++ b/packages/shared-python/shared/services/webhook/dispatcher.py @@ -29,10 +29,10 @@ from shared.models.database.job import Job from shared.models.database.webhook import WebhookEvent, WebhookEventStatus from shared.models.database.webhook_log import WebhookLog -from shared.services.webhook.pinned_outbound_http import ( +from shared.utils.pinned_outbound_http import ( send_pinned_outbound_request, ) -from shared.services.webhook.outbound_url_validator import ( +from shared.utils.outbound_url_validator import ( OutboundURLValidationResult, validate_outbound_url_async, ) diff --git a/packages/shared-python/shared/services/webhook/qstash_publisher.py b/packages/shared-python/shared/services/webhook/qstash_publisher.py index 6bb4d137d..065601a7a 100644 --- a/packages/shared-python/shared/services/webhook/qstash_publisher.py +++ b/packages/shared-python/shared/services/webhook/qstash_publisher.py @@ -22,7 +22,7 @@ from shared.core.config import app_config from shared.core.exceptions.domain_exceptions import QStashServiceException from shared.models.database.webhook import WebhookEventStatus -from shared.services.webhook.outbound_url_validator import validate_outbound_url +from shared.utils.outbound_url_validator import validate_outbound_url class QStashWebhookPublisher: diff --git a/packages/shared-python/shared/services/webhook/outbound_url_validator.py b/packages/shared-python/shared/utils/outbound_url_validator.py similarity index 100% rename from packages/shared-python/shared/services/webhook/outbound_url_validator.py rename to packages/shared-python/shared/utils/outbound_url_validator.py diff --git a/packages/shared-python/shared/services/webhook/pinned_outbound_http.py b/packages/shared-python/shared/utils/pinned_outbound_http.py similarity index 100% rename from packages/shared-python/shared/services/webhook/pinned_outbound_http.py rename to packages/shared-python/shared/utils/pinned_outbound_http.py From 4c13370c0c53b2ec976153f2d1468393d5e4ef69 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Tue, 5 May 2026 10:14:41 +0000 Subject: [PATCH 06/32] refactor: remove database health routes --- apps/api/app/api/v1/health.py | 18 ----- .../contract/test_database_health_contract.py | 57 ---------------- .../shared-python/shared/core/database.py | 68 ------------------- 3 files changed, 143 deletions(-) delete mode 100644 apps/api/tests/contract/test_database_health_contract.py diff --git a/apps/api/app/api/v1/health.py b/apps/api/app/api/v1/health.py index 43837721f..d42aed21e 100644 --- a/apps/api/app/api/v1/health.py +++ b/apps/api/app/api/v1/health.py @@ -4,22 +4,4 @@ from fastapi import APIRouter -from shared.core.database import ( - get_database_performance, - prewarm_connection_pool, -) - router = APIRouter() - - -@router.get("/database/performance") -async def get_database_performance_stats(): - """Return database performance statistics.""" - return await get_database_performance() - - -@router.post("/database/prewarm") -async def prewarm_database_connections(): - """Prewarm the database connection pool.""" - await prewarm_connection_pool() - return {"message": "Database connection pool prewarming completed"} diff --git a/apps/api/tests/contract/test_database_health_contract.py b/apps/api/tests/contract/test_database_health_contract.py deleted file mode 100644 index 135bb0833..000000000 --- a/apps/api/tests/contract/test_database_health_contract.py +++ /dev/null @@ -1,57 +0,0 @@ -from collections.abc import Callable -from contextlib import AbstractAsyncContextManager -from typing import cast - -import pytest -from httpx import AsyncClient - - -@pytest.mark.asyncio -async def test_should_return_the_database_performance_payload_shape( - api_client_factory: Callable[[], AbstractAsyncContextManager[AsyncClient]], -) -> None: - async with api_client_factory() as api_client: - from shared.core.database import db_performance_monitor - - db_performance_monitor.query_times = [] - db_performance_monitor.connection_usage = [] - db_performance_monitor.error_count = 0 - db_performance_monitor.record_query_time(3.2) - db_performance_monitor.record_query_time(6.8) - db_performance_monitor.record_connection_usage( - { - "checked_out": 1, - "checked_in": 2, - "overflow": 0, - } - ) - - response = await api_client.get("/api/v1/health/database/performance") - - assert response.status_code == 200 - - response_json = cast(dict[str, object], response.json()) - query_stats = cast(dict[str, object], response_json["query_stats"]) - connection_stats = cast(dict[str, object], response_json["connection_stats"]) - recent_usage = cast(list[dict[str, object]], connection_stats["recent_usage"]) - - assert query_stats["count"] == 2 - assert isinstance(query_stats["avg_time_ms"], float) - assert isinstance(query_stats["min_time_ms"], float) - assert isinstance(query_stats["max_time_ms"], float) - assert isinstance(query_stats["p95_time_ms"], float) - assert connection_stats["total_errors"] == 0 - assert len(recent_usage) == 1 - - -@pytest.mark.asyncio -async def test_should_return_the_database_prewarm_completion_message( - api_client_factory: Callable[[], AbstractAsyncContextManager[AsyncClient]], -) -> None: - async with api_client_factory() as api_client: - response = await api_client.post("/api/v1/health/database/prewarm") - - assert response.status_code == 200 - assert response.json() == { - "message": "Database connection pool prewarming completed" - } diff --git a/packages/shared-python/shared/core/database.py b/packages/shared-python/shared/core/database.py index 8087d79f2..cd8a03d6e 100644 --- a/packages/shared-python/shared/core/database.py +++ b/packages/shared-python/shared/core/database.py @@ -2,7 +2,6 @@ import logging import os from contextlib import asynccontextmanager -from datetime import datetime from typing import Any, AsyncGenerator, Awaitable, Callable, TypeVar from sqlalchemy import event, text @@ -257,73 +256,6 @@ async def _warm_connection(): logger.debug(f"Connection warming failed: {e}") -# Database performance monitoring. -class DatabasePerformanceMonitor: - """Database performance monitor.""" - - def __init__(self): - self.query_times = [] - self.connection_usage = [] - self.error_count = 0 - - def record_query_time(self, query_time_ms: float): - """Record query latency.""" - self.query_times.append(query_time_ms) - # Keep only the most recent 1000 query samples. - if len(self.query_times) > 1000: - self.query_times = self.query_times[-1000:] - - def record_connection_usage(self, pool_status: dict): - """Record connection-pool usage.""" - self.connection_usage.append( - { - "timestamp": datetime.now().isoformat(), - "checked_out": pool_status.get("checked_out", 0), - "checked_in": pool_status.get("checked_in", 0), - "overflow": pool_status.get("overflow", 0), - } - ) - # Keep only the most recent 100 samples. - if len(self.connection_usage) > 100: - self.connection_usage = self.connection_usage[-100:] - - def record_error(self): - """Record an error occurrence.""" - self.error_count += 1 - - def get_performance_stats(self) -> dict: - """Return collected performance statistics.""" - if not self.query_times: - return {"error": "No query data available"} - - return { - "query_stats": { - "count": len(self.query_times), - "avg_time_ms": round(sum(self.query_times) / len(self.query_times), 2), - "min_time_ms": round(min(self.query_times), 2), - "max_time_ms": round(max(self.query_times), 2), - "p95_time_ms": round( - sorted(self.query_times)[int(len(self.query_times) * 0.95)], 2 - ), - }, - "connection_stats": { - "recent_usage": ( - self.connection_usage[-10:] if self.connection_usage else [] - ), - "total_errors": self.error_count, - }, - } - - -# Shared performance-monitor instance. -db_performance_monitor = DatabasePerformanceMonitor() - - -async def get_database_performance() -> dict: - """Return database performance statistics.""" - return db_performance_monitor.get_performance_stats() - - async def safe_dispose_engine(db_engine: AsyncEngine) -> None: """ Close a database engine safely. From 73ca3cca66754a8aac72e83aeaa68184a1ea787c Mon Sep 17 00:00:00 2001 From: suguanYang Date: Tue, 5 May 2026 10:43:59 +0000 Subject: [PATCH 07/32] ci: add codeql scanning workflow --- .github/workflows/codeql.yml | 45 +++++++++++++++++++++++++ apps/api/scripts/bootstrap_local_dev.py | 11 +++--- 2 files changed, 51 insertions(+), 5 deletions(-) create mode 100644 .github/workflows/codeql.yml diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml new file mode 100644 index 000000000..b3a595948 --- /dev/null +++ b/.github/workflows/codeql.yml @@ -0,0 +1,45 @@ +name: CodeQL + +on: + pull_request: + branches: + - main + - staging + +permissions: + contents: read + security-events: write + actions: read + +concurrency: + group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} + cancel-in-progress: true + +jobs: + analyze: + name: Analyze + runs-on: ubuntu-latest + timeout-minutes: 30 + + steps: + - name: Checkout code + uses: actions/checkout@v6 + with: + persist-credentials: false + + - name: Initialize CodeQL + uses: github/codeql-action/init@v4 + with: + languages: python + queries: security-extended,security-and-quality + + - name: Set up Python + uses: actions/setup-python@v6 + with: + python-version: "3.11" + + - name: Autobuild + uses: github/codeql-action/autobuild@v4 + + - name: Perform CodeQL analysis + uses: github/codeql-action/analyze@v4 diff --git a/apps/api/scripts/bootstrap_local_dev.py b/apps/api/scripts/bootstrap_local_dev.py index e33e895a9..7fbea96b6 100644 --- a/apps/api/scripts/bootstrap_local_dev.py +++ b/apps/api/scripts/bootstrap_local_dev.py @@ -45,11 +45,12 @@ async def _run(mode: str) -> int: def _print_profile() -> None: - print("user_id=local-dev-user") - print("name=Local Development User") - print("email=local-dev-user@knowhere.local") - print("tier=tier_5") - print("local_developer_key_seeded=true") + profile = LocalDevelopmentBootstrapService.get_local_developer_auth_profile() + print(f"user_id={profile['user_id']}") + print(f"name={profile['name']}") + print(f"email={profile['email']}") + print(f"tier={profile['tier']}") + print(f"api_key={profile['api_key']}") def main() -> int: From 20214a760bce230b309be02b95e2c7b06a53548e Mon Sep 17 00:00:00 2001 From: suguanYang Date: Tue, 5 May 2026 14:01:14 +0000 Subject: [PATCH 08/32] fix: pin outbound url downloads --- .../services/storage/sync_storage_service.py | 35 ++- .../contract/test_url_upload_contract.py | 72 +++++ .../services/storage/file_upload_service.py | 66 ++--- .../tests/utils/test_pinned_outbound_http.py | 132 +++++++++ .../shared/utils/pinned_outbound_http.py | 250 +++++++++++++++++- .../shared/utils/url_security.py | 69 ++++- 6 files changed, 545 insertions(+), 79 deletions(-) create mode 100644 packages/shared-python/shared/tests/utils/test_pinned_outbound_http.py diff --git a/apps/worker/app/services/storage/sync_storage_service.py b/apps/worker/app/services/storage/sync_storage_service.py index 800ee5207..d378c9451 100644 --- a/apps/worker/app/services/storage/sync_storage_service.py +++ b/apps/worker/app/services/storage/sync_storage_service.py @@ -6,16 +6,15 @@ import os import tempfile -import uuid as _uuid from typing import Any, Dict, Optional -import requests from loguru import logger from shared.core.config import settings from shared.core.config.storage import get_cached_storage_adapter from shared.core.exceptions.domain_exceptions import StorageServiceException -from shared.utils.url_security import validate_public_http_url +from shared.utils.pinned_outbound_http import download_pinned_outbound_file +from shared.utils.url_security import validate_public_http_url_and_resolve_ip def get_storage_adapter(): @@ -109,25 +108,23 @@ def upload_zip_result(job_id: str, zip_file_path: str) -> str: def download_file_from_url(file_url: str) -> str: - """Download file from URL to temp directory using requests (sync, gevent-compatible).""" - validate_public_http_url(file_url, field="source_url") - - temp_dir = getattr(settings, "TMP_PATH", "/tmp") - os.makedirs(temp_dir, exist_ok=True) - temp_filename = f"temp_{_uuid.uuid4().hex}" - temp_file_path = os.path.join(temp_dir, temp_filename) - + """Download a URL file through SSRF validation and IP pinning.""" + temp_file_path = "" try: - response = requests.get( + validation = validate_public_http_url_and_resolve_ip( file_url, - timeout=300, - stream=True, - headers={"User-Agent": "Knowhere-FileDownloader/1.0"}, + field="source_url", + ) + temp_dir = getattr(settings, "TMP_PATH", "/tmp") + os.makedirs(temp_dir, exist_ok=True) + download_result = download_pinned_outbound_file( + url=validation.url, + pinned_ip=validation.validated_ip, + timeout_seconds=300, + user_agent="Knowhere-FileDownloader/1.0", + temp_dir=temp_dir, ) - response.raise_for_status() - with open(temp_file_path, "wb") as f: - for chunk in response.iter_content(chunk_size=65536): - f.write(chunk) + temp_file_path = download_result.temp_file_path return temp_file_path except Exception as e: if os.path.exists(temp_file_path): diff --git a/apps/worker/tests/contract/test_url_upload_contract.py b/apps/worker/tests/contract/test_url_upload_contract.py index 20de8c376..6ccd4e945 100644 --- a/apps/worker/tests/contract/test_url_upload_contract.py +++ b/apps/worker/tests/contract/test_url_upload_contract.py @@ -3,6 +3,7 @@ import os import socket from pathlib import Path +from types import SimpleNamespace from typing import Any from uuid import uuid4 @@ -140,3 +141,74 @@ def resolve_public_address( assert job_row["status"] == "waiting-file" assert job_row["source_type"] == "url" assert job_row["s3_key"] == s3_key + + +def test_should_download_a_url_file_through_a_pinned_public_ip( + worker_contract_environment: None, + monkeypatch: MonkeyPatch, + tmp_path: Path, +) -> None: + import app.services.storage.sync_storage_service as sync_storage_service + + source_url = "https://example.test/files/contract-source.pdf" + pinned_ip = "93.184.216.34" + validation_calls: list[tuple[str, str]] = [] + download_calls: list[dict[str, object]] = [] + + def fake_validate_public_http_url_and_resolve_ip( + url: str, + field: str = "url", + ) -> SimpleNamespace: + validation_calls.append((url, field)) + return SimpleNamespace(url=url, validated_ip=pinned_ip) + + def fake_download_pinned_outbound_file( + *, + url: str, + pinned_ip: str, + timeout_seconds: float, + user_agent: str, + temp_dir: str | None = None, + field: str = "source_url", + ) -> SimpleNamespace: + download_calls.append( + { + "url": url, + "pinned_ip": pinned_ip, + "timeout_seconds": timeout_seconds, + "user_agent": user_agent, + "temp_dir": temp_dir, + "field": field, + } + ) + temp_file_path = Path(temp_dir or tmp_path) / "downloaded-contract-source.pdf" + temp_file_path.write_bytes(b"pdf") + return SimpleNamespace(status=200, temp_file_path=str(temp_file_path)) + + monkeypatch.setattr( + sync_storage_service, + "validate_public_http_url_and_resolve_ip", + fake_validate_public_http_url_and_resolve_ip, + ) + monkeypatch.setattr( + sync_storage_service, + "download_pinned_outbound_file", + fake_download_pinned_outbound_file, + ) + monkeypatch.setattr(sync_storage_service.settings, "TMP_PATH", str(tmp_path)) + + downloaded_path = sync_storage_service.download_file_from_url(source_url) + + assert validation_calls == [(source_url, "source_url")] + assert download_calls == [ + { + "url": source_url, + "pinned_ip": pinned_ip, + "timeout_seconds": 300, + "user_agent": "Knowhere-FileDownloader/1.0", + "temp_dir": str(tmp_path), + "field": "source_url", + } + ] + assert downloaded_path == str(tmp_path / "downloaded-contract-source.pdf") + assert Path(downloaded_path).read_bytes() == b"pdf" diff --git a/packages/shared-python/shared/services/storage/file_upload_service.py b/packages/shared-python/shared/services/storage/file_upload_service.py index d7b69f55e..e6398a493 100644 --- a/packages/shared-python/shared/services/storage/file_upload_service.py +++ b/packages/shared-python/shared/services/storage/file_upload_service.py @@ -3,10 +3,8 @@ import asyncio import json import os -import uuid from typing import Any, Dict, Optional -import aiohttp from loguru import logger from shared.core.config import settings @@ -14,7 +12,12 @@ KnowhereException, StorageServiceException, ) -from shared.utils.url_security import validate_public_http_url +from shared.utils.pinned_outbound_http import ( + download_pinned_outbound_file_async, +) +from shared.utils.url_security import ( + validate_public_http_url_and_resolve_ip_async, +) class FileUploadService: @@ -427,56 +430,27 @@ def _download(): async def _download_file_from_url(self, file_url: str) -> str: """Download a file from a URL into a temporary directory.""" - validate_public_http_url(file_url, field="source_url") - - temp_dir = getattr(settings, "TMP_PATH", "/tmp") - os.makedirs(temp_dir, exist_ok=True) - - # Generate a temporary filename. - temp_filename = f"temp_{uuid.uuid4().hex}" - temp_file_path = os.path.join(temp_dir, temp_filename) - + temp_file_path = "" try: - # Configure aiohttp for efficient large-file downloads. - timeout = aiohttp.ClientTimeout( - total=300, connect=30 - ) # 5-minute total timeout, 30-second connect timeout. - connector = aiohttp.TCPConnector( - limit=100, # Total connection pool size. - limit_per_host=30, # Connection limit per host. - ttl_dns_cache=300, # Cache DNS for 5 minutes. - use_dns_cache=True, + validation = await validate_public_http_url_and_resolve_ip_async( + file_url, + field="source_url", ) - - async with aiohttp.ClientSession( - timeout=timeout, - connector=connector, - headers={"User-Agent": "Knowhere-FileDownloader/1.0"}, - ) as session: - async with session.get(file_url) as response: - if response.status != 200: - raise StorageServiceException( - internal_message=( - f"Download failed with status code: {response.status}" - ), - operation="download_from_url", - ) - - # Use a larger chunk size to improve download throughput. - with open(temp_file_path, "wb") as f: - async for chunk in response.content.iter_chunked( - 65536 - ): # 64 KB chunks. - f.write(chunk) - + temp_dir = getattr(settings, "TMP_PATH", "/tmp") + os.makedirs(temp_dir, exist_ok=True) + download_result = await download_pinned_outbound_file_async( + url=validation.url, + pinned_ip=validation.validated_ip, + timeout_seconds=300, + user_agent="Knowhere-FileDownloader/1.0", + temp_dir=temp_dir, + ) + temp_file_path = download_result.temp_file_path return temp_file_path except KnowhereException: - if os.path.exists(temp_file_path): - os.remove(temp_file_path) raise except Exception as e: - # Clean up the temporary file on failure. if os.path.exists(temp_file_path): os.remove(temp_file_path) raise StorageServiceException( diff --git a/packages/shared-python/shared/tests/utils/test_pinned_outbound_http.py b/packages/shared-python/shared/tests/utils/test_pinned_outbound_http.py new file mode 100644 index 000000000..e74d02c02 --- /dev/null +++ b/packages/shared-python/shared/tests/utils/test_pinned_outbound_http.py @@ -0,0 +1,132 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Any + +import pytest + +from shared.core.exceptions.domain_exceptions import ValidationException +from shared.utils import pinned_outbound_http + + +class _RedirectResponse: + status = 302 + + def stream(self, chunk_size: int) -> list[bytes]: + return [] + + def release_conn(self) -> None: + return None + + def close(self) -> None: + return None + + +class _RedirectConnectionPool: + def __init__(self, *args: object, **kwargs: object) -> None: + self.args = args + self.kwargs = kwargs + + def urlopen(self, *args: object, **kwargs: object) -> _RedirectResponse: + assert kwargs["redirect"] is False + return _RedirectResponse() + + +class _SuccessResponse: + status = 200 + + def __init__(self) -> None: + self.is_released = False + self.is_closed = False + + def stream(self, chunk_size: int) -> list[bytes]: + return [b"pdf", b""] + + def release_conn(self) -> None: + self.is_released = True + + def close(self) -> None: + self.is_closed = True + + +class _SuccessConnectionPool: + calls: list[dict[str, Any]] = [] + + def __init__(self, *args: object, **kwargs: object) -> None: + self.args = args + self.kwargs = kwargs + + def urlopen(self, *args: object, **kwargs: object) -> _SuccessResponse: + self.calls.append( + { + "init_args": self.args, + "init_kwargs": self.kwargs, + "urlopen_args": args, + "urlopen_kwargs": kwargs, + } + ) + return _SuccessResponse() + + +def test_should_block_redirect_responses_and_remove_partial_download( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + monkeypatch.setattr( + pinned_outbound_http, + "PinnedHTTPConnectionPool", + _RedirectConnectionPool, + ) + + with pytest.raises(ValidationException): + pinned_outbound_http.download_pinned_outbound_file( + url="http://example.test/file.pdf", + pinned_ip="93.184.216.34", + timeout_seconds=300, + user_agent="Knowhere-FileDownloader/1.0", + temp_dir=str(tmp_path), + ) + + assert list(tmp_path.iterdir()) == [] + + +def test_should_request_public_url_through_the_pinned_http_pool( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + _SuccessConnectionPool.calls = [] + monkeypatch.setattr( + pinned_outbound_http, + "PinnedHTTPConnectionPool", + _SuccessConnectionPool, + ) + + result = pinned_outbound_http.download_pinned_outbound_file( + url="http://example.test:8080/files/source.pdf?download=1", + pinned_ip="93.184.216.34", + timeout_seconds=300, + user_agent="Knowhere-FileDownloader/1.0", + temp_dir=str(tmp_path), + ) + + assert Path(result.temp_file_path).read_bytes() == b"pdf" + + call = _SuccessConnectionPool.calls[0] + init_kwargs = call["init_kwargs"] + urlopen_kwargs = call["urlopen_kwargs"] + retry_config = init_kwargs["retries"] + timeout = urlopen_kwargs["timeout"] + + assert call["init_args"] == ("example.test", 8080) + assert init_kwargs["pinned_ip"] == "93.184.216.34" + assert retry_config.total == 0 + assert retry_config.redirect == 0 + assert call["urlopen_args"] == ("GET", "/files/source.pdf?download=1") + assert timeout.connect_timeout == 300 + assert timeout.read_timeout == 300 + assert urlopen_kwargs["preload_content"] is False + assert urlopen_kwargs["redirect"] is False + assert urlopen_kwargs["headers"] == { + "User-Agent": "Knowhere-FileDownloader/1.0", + "Host": "example.test:8080", + } diff --git a/packages/shared-python/shared/utils/pinned_outbound_http.py b/packages/shared-python/shared/utils/pinned_outbound_http.py index 6b4db161a..808b219c8 100644 --- a/packages/shared-python/shared/utils/pinned_outbound_http.py +++ b/packages/shared-python/shared/utils/pinned_outbound_http.py @@ -1,20 +1,24 @@ -""" -Pinned outbound HTTP helpers. - -Shared infrastructure for outbound requests that must connect to a -pre-validated public IP address and block redirect-based SSRF. -""" +"""Pinned outbound HTTP helpers.""" from __future__ import annotations +import os import socket +import tempfile from collections.abc import Mapping from dataclasses import dataclass from typing import Any +from urllib.parse import urlsplit import aiohttp from aiohttp.abc import AbstractResolver +from urllib3 import Retry +from urllib3.connection import HTTPConnection, HTTPSConnection +from urllib3.connectionpool import HTTPConnectionPool, HTTPSConnectionPool +from urllib3.util import Timeout +from urllib3.util.connection import create_connection +from shared.core.exceptions.domain_exceptions import ValidationException from shared.utils.url_security import SafePublicHTTPURL @@ -25,6 +29,14 @@ class PinnedOutboundResponse: status: int +@dataclass(frozen=True) +class PinnedDownloadResult: + """Metadata for a pinned HTTP download to a temporary file.""" + + status: int + temp_file_path: str + + class PinnedIPResolver(AbstractResolver): """Resolver that always returns the supplied pinned IP address.""" @@ -55,6 +67,90 @@ async def close(self) -> None: pass +class PinnedHTTPConnection(HTTPConnection): + """HTTP connection that resolves the hostname to a pinned IP address.""" + + def __init__(self, *args: Any, pinned_ip: str, **kwargs: Any) -> None: + self._pinned_ip = pinned_ip + super().__init__(*args, **kwargs) + + def _new_conn(self) -> socket.socket: + return create_connection( + (self._pinned_ip, self.port), + self.timeout, + source_address=self.source_address, + socket_options=self.socket_options, + ) + + +class PinnedHTTPSConnection(HTTPSConnection): + """HTTPS connection that resolves the hostname to a pinned IP address.""" + + def __init__(self, *args: Any, pinned_ip: str, **kwargs: Any) -> None: + self._pinned_ip = pinned_ip + super().__init__(*args, **kwargs) + + def _new_conn(self) -> socket.socket: + return create_connection( + (self._pinned_ip, self.port), + self.timeout, + source_address=self.source_address, + socket_options=self.socket_options, + ) + + +class PinnedHTTPConnectionPool(HTTPConnectionPool): + """HTTP pool that pins DNS resolution to a fixed IP address.""" + + ConnectionCls = PinnedHTTPConnection # pyright: ignore[reportAssignmentType] + + def __init__(self, *args: Any, pinned_ip: str, **kwargs: Any) -> None: + kwargs["pinned_ip"] = pinned_ip + super().__init__(*args, **kwargs) + + +class PinnedHTTPSConnectionPool(HTTPSConnectionPool): + """HTTPS pool that pins DNS resolution to a fixed IP address.""" + + ConnectionCls = PinnedHTTPSConnection # pyright: ignore[reportAssignmentType] + + def __init__(self, *args: Any, pinned_ip: str, **kwargs: Any) -> None: + kwargs["pinned_ip"] = pinned_ip + super().__init__(*args, **kwargs) + + +def _build_host_header(parsed_url: Any) -> str: + hostname = parsed_url.hostname + if not hostname: + return "" + + if ":" in hostname and not hostname.startswith("["): + formatted_host = f"[{hostname}]" + else: + formatted_host = hostname + + if parsed_url.port is not None: + return f"{formatted_host}:{parsed_url.port}" + return formatted_host + + +def _validate_download_url(url: str, field: str) -> SafePublicHTTPURL: + parsed_url = urlsplit(url) + if parsed_url.scheme not in {"http", "https"}: + raise ValidationException( + user_message="Invalid URL", + violations=[ + {"field": field, "description": "URL must use http or https"} + ], + ) + if not parsed_url.hostname: + raise ValidationException( + user_message="Invalid URL", + violations=[{"field": field, "description": "URL must include a hostname"}], + ) + return SafePublicHTTPURL(url) + + async def send_pinned_outbound_request( *, method: str, @@ -64,12 +160,8 @@ async def send_pinned_outbound_request( headers: Mapping[str, str] | None = None, json_body: Any | None = None, ) -> PinnedOutboundResponse: - """ - Send an outbound HTTP request through a resolver pinned to a validated IP. - - Redirects are always blocked to prevent redirect-based SSRF. - """ - validated_url = SafePublicHTTPURL(url) + """Send an outbound HTTP request through a resolver pinned to a validated IP.""" + validated_url = _validate_download_url(url, "url") connector = aiohttp.TCPConnector(resolver=PinnedIPResolver(pinned_ip)) timeout = aiohttp.ClientTimeout(total=timeout_seconds) @@ -85,3 +177,137 @@ async def send_pinned_outbound_request( allow_redirects=False, ) as response: return PinnedOutboundResponse(status=response.status) + + +async def download_pinned_outbound_file_async( + *, + url: str, + pinned_ip: str, + timeout_seconds: float, + user_agent: str, + temp_dir: str | None = None, + field: str = "source_url", +) -> PinnedDownloadResult: + """Download a file with aiohttp while pinning DNS to a validated IP.""" + validated_url = _validate_download_url(url, field) + + temp_file_descriptor, temp_file_path = tempfile.mkstemp(dir=temp_dir) + os.close(temp_file_descriptor) + + connector = aiohttp.TCPConnector(resolver=PinnedIPResolver(pinned_ip)) + timeout = aiohttp.ClientTimeout(total=timeout_seconds) + + try: + async with aiohttp.ClientSession( + connector=connector, + timeout=timeout, + ) as session: + async with session.get( + validated_url, + allow_redirects=False, + headers={"User-Agent": user_agent}, + ) as response: + if not 200 <= response.status < 300: + raise ValidationException( + user_message="Invalid URL", + violations=[ + { + "field": field, + "description": f"URL request failed with status {response.status}", + } + ], + ) + + with open(temp_file_path, "wb") as output_file: + async for chunk in response.content.iter_chunked(65536): + if chunk: + output_file.write(chunk) + + return PinnedDownloadResult( + status=response.status, + temp_file_path=temp_file_path, + ) + except Exception: + if os.path.exists(temp_file_path): + os.remove(temp_file_path) + raise + + +def download_pinned_outbound_file( + *, + url: str, + pinned_ip: str, + timeout_seconds: float, + user_agent: str, + temp_dir: str | None = None, + field: str = "source_url", +) -> PinnedDownloadResult: + """ + Download a file through a resolver pinned to a validated IP. + + Redirects are blocked to prevent redirect-based SSRF. The request connects + to the pre-validated IP while preserving the original Host header. + """ + validated_url = _validate_download_url(url, field) + parsed_url = urlsplit(validated_url) + + temp_file_descriptor, temp_file_path = tempfile.mkstemp(dir=temp_dir) + os.close(temp_file_descriptor) + + try: + connection_pool: HTTPConnectionPool + request_path = parsed_url.path or "/" + if parsed_url.query: + request_path = f"{request_path}?{parsed_url.query}" + + if parsed_url.scheme == "https": + connection_pool = PinnedHTTPSConnectionPool( + parsed_url.hostname, + parsed_url.port or 443, + pinned_ip=pinned_ip, + retries=Retry(total=0, redirect=False), + ) + else: + connection_pool = PinnedHTTPConnectionPool( + parsed_url.hostname, + parsed_url.port or 80, + pinned_ip=pinned_ip, + retries=Retry(total=0, redirect=False), + ) + + response = connection_pool.urlopen( + "GET", + request_path, + timeout=Timeout.from_float(timeout_seconds), + preload_content=False, + redirect=False, + headers={"User-Agent": user_agent, "Host": _build_host_header(parsed_url)}, + ) + try: + if not 200 <= response.status < 300: + raise ValidationException( + user_message="Invalid URL", + violations=[ + { + "field": field, + "description": f"URL request failed with status {response.status}", + } + ], + ) + + with open(temp_file_path, "wb") as output_file: + for chunk in response.stream(65536): + if chunk: + output_file.write(chunk) + finally: + response.release_conn() + response.close() + + return PinnedDownloadResult( + status=response.status, + temp_file_path=temp_file_path, + ) + except Exception: + if os.path.exists(temp_file_path): + os.remove(temp_file_path) + raise diff --git a/packages/shared-python/shared/utils/url_security.py b/packages/shared-python/shared/utils/url_security.py index 34154417b..0b76c06be 100644 --- a/packages/shared-python/shared/utils/url_security.py +++ b/packages/shared-python/shared/utils/url_security.py @@ -1,6 +1,7 @@ import asyncio import ipaddress import socket +from dataclasses import dataclass from typing import cast from urllib.parse import urljoin, urlparse @@ -39,6 +40,14 @@ class SafePublicHTTPURL(str): """A URL string that has passed public HTTP SSRF validation.""" +@dataclass(frozen=True) +class PublicHTTPURLValidationResult: + """Validated public HTTP URL and its pinned public IP address.""" + + url: SafePublicHTTPURL + validated_ip: str + + def validate_public_http_url(url: str, field: str = "url") -> None: """Reject URL inputs that could target internal networks or local services.""" try: @@ -72,6 +81,62 @@ def get_safe_public_http_url(url: str, field: str = "url") -> SafePublicHTTPURL: return SafePublicHTTPURL(url) +def validate_public_http_url_and_resolve_ip( + url: str, + field: str = "url", +) -> PublicHTTPURLValidationResult: + """Validate a public HTTP URL and return the IP selected during validation.""" + try: + validated_ip = _validate_public_http_url(url) + return PublicHTTPURLValidationResult( + url=SafePublicHTTPURL(url), + validated_ip=validated_ip, + ) + except UnsupportedURLSchemeError as exc: + raise _build_url_validation_error(field, "URL must use http or https") from exc + except MissingURLHostnameError as exc: + raise _build_url_validation_error(field, "URL must include a hostname") from exc + except HostnameResolutionError as exc: + raise _build_url_validation_error(field, "URL hostname could not be resolved") from exc + except InvalidResolvedAddressError as exc: + raise _build_url_validation_error(field, "URL resolved to an invalid IP") from exc + except HostnameNotAllowedError as exc: + raise _build_url_validation_error(field, "URL host is not allowed") from exc + + +async def validate_public_http_url_and_resolve_ip_async( + url: str, + field: str = "url", +) -> PublicHTTPURLValidationResult: + """Validate a public HTTP URL asynchronously and return the selected IP.""" + try: + parsed_url = urlparse(url) + if parsed_url.scheme not in {"http", "https"}: + raise UnsupportedURLSchemeError( + f"Unsupported URL scheme: {parsed_url.scheme}" + ) + + hostname = parsed_url.hostname + if not hostname: + raise MissingURLHostnameError("URL must include a hostname") + + validated_ip = await resolve_public_hostname_async(hostname) + return PublicHTTPURLValidationResult( + url=SafePublicHTTPURL(url), + validated_ip=validated_ip, + ) + except UnsupportedURLSchemeError as exc: + raise _build_url_validation_error(field, "URL must use http or https") from exc + except MissingURLHostnameError as exc: + raise _build_url_validation_error(field, "URL must include a hostname") from exc + except HostnameResolutionError as exc: + raise _build_url_validation_error(field, "URL hostname could not be resolved") from exc + except InvalidResolvedAddressError as exc: + raise _build_url_validation_error(field, "URL resolved to an invalid IP") from exc + except HostnameNotAllowedError as exc: + raise _build_url_validation_error(field, "URL host is not allowed") from exc + + def get_safe_redirect_url(url: str, redirect_url: str) -> str: """Resolve and validate an HTTP redirect target for internal network callers.""" resolved_url = urljoin(url, redirect_url) @@ -92,7 +157,7 @@ def resolve_public_hostname(hostname: str) -> str: return _select_public_ip_address(hostname, address_infos) -def _validate_public_http_url(url: str) -> None: +def _validate_public_http_url(url: str) -> str: parsed_url = urlparse(url) if parsed_url.scheme not in {"http", "https"}: raise UnsupportedURLSchemeError(f"Unsupported URL scheme: {parsed_url.scheme}") @@ -101,7 +166,7 @@ def _validate_public_http_url(url: str) -> None: if not hostname: raise MissingURLHostnameError("URL must include a hostname") - resolve_public_hostname(hostname) + return resolve_public_hostname(hostname) async def resolve_public_hostname_async(hostname: str) -> str: From e9d3a396bc1374e7b39f67fcfd64621726606018 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 01:34:55 +0000 Subject: [PATCH 09/32] fix: simplify api key hashing --- .../f6a7b8c9d0e1_key_api_key_hashes.py | 25 ------------ apps/api/app/core/dependencies.py | 7 +++- apps/api/app/services/auth/api_key_service.py | 27 ++++++------- .../guest/guest_registration_service.py | 7 ++-- .../app/services/rate_limit/dependencies.py | 40 ++++++++++++------- apps/api/scripts/init_user.py | 17 ++------ .../scripts/local_dev_bootstrap_service.py | 10 +---- .../tests/contract/test_api_key_contract.py | 34 ++++++++++++++++ apps/api/tests/support/contract_database.py | 4 +- .../shared/tests/utils/test_api_keys.py | 34 ++++++++++++++++ .../shared/utils/api_key_hashing.py | 13 ------ .../shared-python/shared/utils/api_keys.py | 29 ++++++++++++++ 12 files changed, 151 insertions(+), 96 deletions(-) delete mode 100644 apps/api/alembic/versions/f6a7b8c9d0e1_key_api_key_hashes.py create mode 100644 packages/shared-python/shared/tests/utils/test_api_keys.py delete mode 100644 packages/shared-python/shared/utils/api_key_hashing.py create mode 100644 packages/shared-python/shared/utils/api_keys.py diff --git a/apps/api/alembic/versions/f6a7b8c9d0e1_key_api_key_hashes.py b/apps/api/alembic/versions/f6a7b8c9d0e1_key_api_key_hashes.py deleted file mode 100644 index cd63a6734..000000000 --- a/apps/api/alembic/versions/f6a7b8c9d0e1_key_api_key_hashes.py +++ /dev/null @@ -1,25 +0,0 @@ -"""key api key hashes - -Revision ID: f6a7b8c9d0e1 -Revises: e5f6a7b8c9d0 -Create Date: 2026-05-01 05:55:00.000000 - -""" - -from typing import Sequence, Union - -from alembic import op - - -revision: str = "f6a7b8c9d0e1" -down_revision: Union[str, Sequence[str], None] = "e5f6a7b8c9d0" -branch_labels: Union[str, Sequence[str], None] = None -depends_on: Union[str, Sequence[str], None] = None - - -def upgrade() -> None: - op.execute("UPDATE api_keys SET is_active = false") - - -def downgrade() -> None: - pass diff --git a/apps/api/app/core/dependencies.py b/apps/api/app/core/dependencies.py index a234bb92e..89d2458f8 100644 --- a/apps/api/app/core/dependencies.py +++ b/apps/api/app/core/dependencies.py @@ -17,7 +17,7 @@ AuthException, PermissionDeniedException, ) -from shared.utils.api_key_hashing import hash_api_key +from shared.utils.api_keys import hash_api_key # Standard JWKS endpoint path (fixed, following OpenID Connect convention) JWKS_ENDPOINT_PATH = "/api/auth/jwks" @@ -188,8 +188,9 @@ async def get_current_user_id( # Check identity cache first — skip DB on cache hit api_key_hash = hash_api_key(token) try: + redis_service = redis_pool_manager.get_redis_service() cached = await identity_cache.get_cached_identity( - redis_pool_manager.get_redis_service(), + redis_service, identity_cache._apikey_key(api_key_hash), ) if cached is not None: @@ -199,6 +200,7 @@ async def get_current_user_id( request.state.cached_user_tier = cached_user_tier request.state.cached_identity_hit = True request.state.user_id = cached_user_id + request.state.api_key_hash = api_key_hash _enforce_guest_api_key_scope(route_path, cached_user_tier) return cached_user_id except PermissionDeniedException: @@ -213,6 +215,7 @@ async def get_current_user_id( request.state.cached_user_tier = identity.user_tier request.state.cached_identity_hit = False request.state.user_id = identity.user_id + request.state.api_key_hash = identity.key_hash _enforce_guest_api_key_scope(route_path, identity.user_tier) return identity.user_id else: diff --git a/apps/api/app/services/auth/api_key_service.py b/apps/api/app/services/auth/api_key_service.py index c9f50d08c..5e81ae592 100644 --- a/apps/api/app/services/auth/api_key_service.py +++ b/apps/api/app/services/auth/api_key_service.py @@ -1,7 +1,6 @@ """API key management service.""" import asyncio -import uuid from dataclasses import dataclass from datetime import datetime from typing import List, Optional @@ -22,7 +21,11 @@ ) from shared.models.database.api_key import APIKey from shared.models.database.user_balance import UserBalance -from shared.utils.api_key_hashing import hash_api_key +from shared.utils.api_keys import ( + generate_api_key, + hash_api_key, + mask_api_key, +) _DEFAULT_USER_TIER: str = "free" @@ -33,6 +36,7 @@ class APIKeyIdentity: user_id: str user_tier: str + key_hash: str class APIKeyService: @@ -41,12 +45,6 @@ class APIKeyService: def __init__(self): self.repository = APIKeyRepository() - def _mask_api_key(self, api_key: str) -> str: - """Mask an API key, exposing only the first 8 and last 4 characters.""" - if not api_key or len(api_key) < 12: - return api_key - return api_key[:8] + "•" * (len(api_key) - 12) + api_key[-4:] - async def create_api_key( self, session: AsyncSession, @@ -83,10 +81,10 @@ async def create_api_key( ], ) - # 3. Generate a secure API key (sk_ + a 32-char UUID without hyphens). - api_key = f"sk_{str(uuid.uuid4()).replace('-', '')}" + # 3. Generate and store a secure API key. + api_key = generate_api_key() key_hash = hash_api_key(api_key) - key_mask = self._mask_api_key(api_key) + key_mask = mask_api_key(api_key) # 4. Store it in the database. api_key_record = APIKey( @@ -129,6 +127,7 @@ async def validate_api_key_identity( return APIKeyIdentity( user_id=user_id, user_tier=user_tier, + key_hash=str(api_key_record.key_hash), ) async def _resolve_user_tier( @@ -234,10 +233,10 @@ async def regenerate_api_key( internal_message="API Key not found or does not belong to user", ) - # 2. Generate a new API key (sk_ + a 32-char UUID without hyphens). - new_api_key = f"sk_{str(uuid.uuid4()).replace('-', '')}" + # 2. Generate a new API key. + new_api_key = generate_api_key() new_key_hash = hash_api_key(new_api_key) - new_key_mask = self._mask_api_key(new_api_key) + new_key_mask = mask_api_key(new_api_key) # 3. Update the database record. from sqlalchemy import update diff --git a/apps/api/app/services/guest/guest_registration_service.py b/apps/api/app/services/guest/guest_registration_service.py index f7e5fcafe..6cd16ee63 100644 --- a/apps/api/app/services/guest/guest_registration_service.py +++ b/apps/api/app/services/guest/guest_registration_service.py @@ -1,7 +1,6 @@ """Guest registration business logic.""" import hashlib -import uuid from datetime import datetime from typing import NoReturn from uuid import uuid4 @@ -26,7 +25,7 @@ GuestRegisterResponse, ) from shared.services.billing.credits_service import CreditsService -from shared.utils.api_key_hashing import hash_api_key +from shared.utils.api_keys import generate_api_key, hash_api_key, mask_api_key _GUEST_TIER: str = "guest" _GUEST_KEY_NAME_PREFIX: str = "guest-device" @@ -157,9 +156,9 @@ async def _create_api_key_without_commit( """ from shared.models.database.api_key import APIKey - api_key = f"sk_{str(uuid.uuid4()).replace('-', '')}" + api_key = generate_api_key() key_hash = hash_api_key(api_key) - key_mask = self._api_key_service._mask_api_key(api_key) + key_mask = mask_api_key(api_key) api_key_record = APIKey( user_id=user_id, diff --git a/apps/api/app/services/rate_limit/dependencies.py b/apps/api/app/services/rate_limit/dependencies.py index c70f6d8c5..8632f1eed 100644 --- a/apps/api/app/services/rate_limit/dependencies.py +++ b/apps/api/app/services/rate_limit/dependencies.py @@ -41,9 +41,9 @@ from shared.core.logging import log_context from shared.core.state_machine.states import JobStatus from shared.models.database.api_key import APIKey -from shared.utils.api_key_hashing import hash_api_key from shared.models.database.job import Job from shared.models.database.user_balance import UserBalance +from shared.utils.api_keys import hash_api_key, is_api_key_token _DEFAULT_TIER: str = "free" _ACTIVE_JOB_STATES: tuple[str, ...] = ( @@ -172,11 +172,11 @@ async def with_current_user( if isinstance(user_tier, str) and stashed_user_id == user_id: if cached_identity_hit is False: token = _extract_bearer_token(request.headers.get("authorization")) - api_key_hash = None - is_api_key_auth = isinstance(token, str) and token.startswith("sk_") - if token is not None and is_api_key_auth: + api_key_hash = getattr(request.state, "api_key_hash", None) + is_api_key_auth = is_api_key_token(token) + if is_api_key_auth and not api_key_hash and isinstance(token, str): api_key_hash = hash_api_key(token) - if is_api_key_auth and api_key_hash: + if is_api_key_auth and isinstance(api_key_hash, str): try: ttl_seconds = await _resolve_apikey_cache_ttl_seconds(api_key_hash) await identity_cache.set_apikey_identity( @@ -194,24 +194,36 @@ async def with_current_user( ) else: token = _extract_bearer_token(request.headers.get("authorization")) - api_key_hash = None - is_api_key_auth = isinstance(token, str) and token.startswith("sk_") - if token is not None and is_api_key_auth: - api_key_hash = hash_api_key(token) + is_api_key_auth = is_api_key_token(token) + api_key_hash = ( + hash_api_key(token) + if is_api_key_auth and isinstance(token, str) + else None + ) cache_key: str = ( identity_cache._apikey_key(api_key_hash) - if is_api_key_auth and api_key_hash + if isinstance(api_key_hash, str) else identity_cache._jwt_key(user_id) ) try: - cached: dict | None = await identity_cache.get_cached_identity( - redis_service, cache_key - ) + if isinstance(api_key_hash, str): + cached = await identity_cache.get_cached_identity( + redis_service, + identity_cache._apikey_key(api_key_hash), + ) + else: + cached = await identity_cache.get_cached_identity( + redis_service, + cache_key, + ) if cached is not None: user_tier = cached.get("user_tier", _DEFAULT_TIER) + if isinstance(api_key_hash, str): + request.state.api_key_hash = api_key_hash else: user_tier = await _resolve_user_tier_from_db(user_id) - if is_api_key_auth and api_key_hash: + api_key_hash = getattr(request.state, "api_key_hash", None) + if is_api_key_auth and isinstance(api_key_hash, str): ttl_seconds = await _resolve_apikey_cache_ttl_seconds(api_key_hash) await identity_cache.set_apikey_identity( redis_service, diff --git a/apps/api/scripts/init_user.py b/apps/api/scripts/init_user.py index 0ae8fa904..c7442fc85 100644 --- a/apps/api/scripts/init_user.py +++ b/apps/api/scripts/init_user.py @@ -3,7 +3,6 @@ import argparse import asyncio import os -import secrets import sys from pathlib import Path from uuid import uuid4 @@ -20,7 +19,7 @@ from shared.models.database.user import User from shared.models.database.user_balance import UserBalance from shared.services.auth.user_table_bootstrap import ensure_better_auth_user_table -from shared.utils.api_key_hashing import hash_api_key +from shared.utils.api_keys import generate_api_key, hash_api_key, mask_api_key _DEFAULT_API_KEY_NAME: str = "standalone-api-key" _DEFAULT_USER_TIER: str = "free" @@ -130,16 +129,6 @@ async def _resolve_key_name( return f"{key_name}-{suffix}" -def _generate_api_key() -> str: - return f"sk_kn_{secrets.token_hex(16)}" - - -def _mask_api_key(api_key: str) -> str: - if len(api_key) < 12: - return api_key - return api_key[:8] + "•" * (len(api_key) - 12) + api_key[-4:] - - def _write_api_key_file(path_value: str, api_key: str) -> Path: output_path = Path(path_value).expanduser() output_path.parent.mkdir(parents=True, exist_ok=True) @@ -159,12 +148,12 @@ async def _create_api_key( user_id: str, key_name: str, ) -> str: - api_key = _generate_api_key() + api_key = generate_api_key() session.add( APIKey( user_id=user_id, key_hash=hash_api_key(api_key), - key_mask=_mask_api_key(api_key), + key_mask=mask_api_key(api_key), name=key_name, enabled_modules=["all"], ) diff --git a/apps/api/scripts/local_dev_bootstrap_service.py b/apps/api/scripts/local_dev_bootstrap_service.py index d4ac0217c..7a98088ab 100644 --- a/apps/api/scripts/local_dev_bootstrap_service.py +++ b/apps/api/scripts/local_dev_bootstrap_service.py @@ -12,7 +12,7 @@ from shared.models.database.user import User from shared.models.database.user_balance import UserBalance from shared.services.auth.user_table_bootstrap import ensure_better_auth_user_table -from shared.utils.api_key_hashing import hash_api_key +from shared.utils.api_keys import hash_api_key, mask_api_key class LocalDevelopmentBootstrapService: @@ -156,7 +156,7 @@ async def _upsert_credits_transaction(self, session: AsyncSession) -> None: async def _upsert_api_key(self, session: AsyncSession) -> None: api_key = await session.get(APIKey, self.LOCAL_DEV_API_KEY_ID) key_hash = hash_api_key(self.LOCAL_DEV_API_KEY) - key_mask = self._mask_api_key(self.LOCAL_DEV_API_KEY) + key_mask = mask_api_key(self.LOCAL_DEV_API_KEY) if api_key is None: session.add( @@ -178,12 +178,6 @@ async def _upsert_api_key(self, session: AsyncSession) -> None: api_key.enabled_modules = ["all"] api_key.is_active = True - @staticmethod - def _mask_api_key(api_key: str) -> str: - if len(api_key) < 12: - return api_key - return api_key[:8] + "•" * (len(api_key) - 12) + api_key[-4:] - @staticmethod def _utc_now() -> datetime: return datetime.now(timezone.utc).replace(tzinfo=None) diff --git a/apps/api/tests/contract/test_api_key_contract.py b/apps/api/tests/contract/test_api_key_contract.py index 21b693fb3..ec55ff866 100644 --- a/apps/api/tests/contract/test_api_key_contract.py +++ b/apps/api/tests/contract/test_api_key_contract.py @@ -6,6 +6,9 @@ import pytest from httpx import AsyncClient +from shared.utils.api_keys import hash_api_key +from tests.support.contract_database import ContractDatabase + @pytest.mark.asyncio async def test_should_revoke_a_created_api_key_through_http_only( @@ -78,6 +81,37 @@ async def test_should_revoke_a_created_api_key_through_http_only( assert "details" not in error +@pytest.mark.asyncio +async def test_should_accept_an_active_sha256_api_key_hash( + api_client_factory: Callable[[], AbstractAsyncContextManager[AsyncClient]], +) -> None: + user_id = f"sha256-user-{uuid4().hex[:12]}" + raw_api_key = f"sk_sha256_{uuid4().hex}" + key_hash = hash_api_key(raw_api_key) + + async with api_client_factory() as api_client: + await ContractDatabase.insert_authenticated_user( + user_id=user_id, + api_key=raw_api_key, + user_tier="tier_5", + ) + await ContractDatabase.execute( + """ + UPDATE api_keys + SET key_hash = :key_hash + WHERE user_id = :user_id + """, + { + "key_hash": key_hash, + "user_id": user_id, + }, + ) + api_client.headers.update({"Authorization": f"Bearer {raw_api_key}"}) + response = await api_client.get("/api/v1/jobs") + + assert response.status_code == 200 + + @pytest.mark.asyncio async def test_should_regenerate_an_api_key_and_invalidate_the_previous_raw_key( developer_api_client_factory: Callable[ diff --git a/apps/api/tests/support/contract_database.py b/apps/api/tests/support/contract_database.py index 28f955e10..1a6d34bbf 100644 --- a/apps/api/tests/support/contract_database.py +++ b/apps/api/tests/support/contract_database.py @@ -9,7 +9,7 @@ from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine from shared.testing.contract_runtime import get_contract_database_url -from shared.utils.api_key_hashing import hash_api_key +from shared.utils.api_keys import hash_api_key, mask_api_key async def _create_contract_engine() -> AsyncEngine: @@ -168,7 +168,7 @@ async def insert_authenticated_user( "id": api_key_id, "user_id": user_id, "key_hash": api_key_hash, - "key_mask": f"{api_key[:8]}...{api_key[-4:]}", + "key_mask": mask_api_key(api_key), "name": f"Contract API Key {user_id}", "enabled_modules": json.dumps(enabled_modules or ["all"]), "is_active": True, diff --git a/packages/shared-python/shared/tests/utils/test_api_keys.py b/packages/shared-python/shared/tests/utils/test_api_keys.py new file mode 100644 index 000000000..c4984474d --- /dev/null +++ b/packages/shared-python/shared/tests/utils/test_api_keys.py @@ -0,0 +1,34 @@ +from shared.utils.api_keys import ( + API_KEY_PREFIX, + generate_api_key, + hash_api_key, + is_api_key_token, + mask_api_key, +) + + +def test_generate_api_key_should_use_api_key_prefix_and_random_secret() -> None: + first_api_key = generate_api_key() + second_api_key = generate_api_key() + + assert first_api_key.startswith(API_KEY_PREFIX) + assert second_api_key.startswith(API_KEY_PREFIX) + assert first_api_key != second_api_key + assert len(first_api_key) > len(API_KEY_PREFIX) + 32 + + +def test_hash_api_key_should_return_deterministic_sha256_lookup_hash() -> None: + api_key = "sk_contract_test_secret" + + assert hash_api_key(api_key) == hash_api_key(api_key) + assert len(hash_api_key(api_key)) == 64 + + +def test_mask_api_key_should_hide_middle_characters() -> None: + assert mask_api_key("sk_1234567890abcdef") == "sk_12345•••••••cdef" + + +def test_is_api_key_token_should_match_only_api_key_prefix() -> None: + assert is_api_key_token("sk_test") is True + assert is_api_key_token("jwt_test") is False + assert is_api_key_token(None) is False diff --git a/packages/shared-python/shared/utils/api_key_hashing.py b/packages/shared-python/shared/utils/api_key_hashing.py deleted file mode 100644 index 46dfe705a..000000000 --- a/packages/shared-python/shared/utils/api_key_hashing.py +++ /dev/null @@ -1,13 +0,0 @@ -"""API key hashing helpers.""" - -import hmac -from hashlib import sha256 - - -def hash_api_key(api_key: str) -> str: - """Return a deterministic keyed digest for API key lookup.""" - from shared.core.config import settings - - secret_key = settings.SECRET_KEY.encode("utf-8") - api_key_bytes = api_key.encode("utf-8") - return hmac.new(secret_key, api_key_bytes, sha256).hexdigest() diff --git a/packages/shared-python/shared/utils/api_keys.py b/packages/shared-python/shared/utils/api_keys.py new file mode 100644 index 000000000..a9665de04 --- /dev/null +++ b/packages/shared-python/shared/utils/api_keys.py @@ -0,0 +1,29 @@ +"""API key generation, masking, and hashing helpers.""" + +from hashlib import sha256 +from secrets import token_urlsafe + +API_KEY_PREFIX = "sk_" +API_KEY_RANDOM_BYTES = 32 + + +def hash_api_key(api_key: str) -> str: + """Return a deterministic SHA-256 digest for API key lookup.""" + return sha256(api_key.encode("utf-8")).hexdigest() + + +def generate_api_key() -> str: + """Generate a new plaintext API key with cryptographic randomness.""" + return f"{API_KEY_PREFIX}{token_urlsafe(API_KEY_RANDOM_BYTES)}" + + +def mask_api_key(api_key: str) -> str: + """Mask an API key, exposing only the first 8 and last 4 characters.""" + if len(api_key) < 12: + return api_key + return api_key[:8] + "•" * (len(api_key) - 12) + api_key[-4:] + + +def is_api_key_token(token: str | None) -> bool: + """Return whether a bearer token has the API-key prefix.""" + return isinstance(token, str) and token.startswith(API_KEY_PREFIX) From fd97cf6784e20f508f79d2767aa7608f64ba702d Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 02:56:00 +0000 Subject: [PATCH 10/32] refactor: simplify api key token narrowing --- .../api/app/services/rate_limit/dependencies.py | 17 ++++++----------- packages/shared-python/shared/utils/api_keys.py | 3 ++- 2 files changed, 8 insertions(+), 12 deletions(-) diff --git a/apps/api/app/services/rate_limit/dependencies.py b/apps/api/app/services/rate_limit/dependencies.py index 8632f1eed..d92ecd8fd 100644 --- a/apps/api/app/services/rate_limit/dependencies.py +++ b/apps/api/app/services/rate_limit/dependencies.py @@ -173,10 +173,9 @@ async def with_current_user( if cached_identity_hit is False: token = _extract_bearer_token(request.headers.get("authorization")) api_key_hash = getattr(request.state, "api_key_hash", None) - is_api_key_auth = is_api_key_token(token) - if is_api_key_auth and not api_key_hash and isinstance(token, str): + if not isinstance(api_key_hash, str) and is_api_key_token(token): api_key_hash = hash_api_key(token) - if is_api_key_auth and isinstance(api_key_hash, str): + if isinstance(api_key_hash, str): try: ttl_seconds = await _resolve_apikey_cache_ttl_seconds(api_key_hash) await identity_cache.set_apikey_identity( @@ -194,12 +193,9 @@ async def with_current_user( ) else: token = _extract_bearer_token(request.headers.get("authorization")) - is_api_key_auth = is_api_key_token(token) - api_key_hash = ( - hash_api_key(token) - if is_api_key_auth and isinstance(token, str) - else None - ) + api_key_hash: str | None = None + if is_api_key_token(token): + api_key_hash = hash_api_key(token) cache_key: str = ( identity_cache._apikey_key(api_key_hash) if isinstance(api_key_hash, str) @@ -222,8 +218,7 @@ async def with_current_user( request.state.api_key_hash = api_key_hash else: user_tier = await _resolve_user_tier_from_db(user_id) - api_key_hash = getattr(request.state, "api_key_hash", None) - if is_api_key_auth and isinstance(api_key_hash, str): + if isinstance(api_key_hash, str): ttl_seconds = await _resolve_apikey_cache_ttl_seconds(api_key_hash) await identity_cache.set_apikey_identity( redis_service, diff --git a/packages/shared-python/shared/utils/api_keys.py b/packages/shared-python/shared/utils/api_keys.py index a9665de04..71b5f8add 100644 --- a/packages/shared-python/shared/utils/api_keys.py +++ b/packages/shared-python/shared/utils/api_keys.py @@ -2,6 +2,7 @@ from hashlib import sha256 from secrets import token_urlsafe +from typing import TypeGuard API_KEY_PREFIX = "sk_" API_KEY_RANDOM_BYTES = 32 @@ -24,6 +25,6 @@ def mask_api_key(api_key: str) -> str: return api_key[:8] + "•" * (len(api_key) - 12) + api_key[-4:] -def is_api_key_token(token: str | None) -> bool: +def is_api_key_token(token: object) -> TypeGuard[str]: """Return whether a bearer token has the API-key prefix.""" return isinstance(token, str) and token.startswith(API_KEY_PREFIX) From 0edceb9cd792507785e197030d42e8eef66cfcf3 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 03:22:28 +0000 Subject: [PATCH 11/32] refactor: separate api key auth cache from rate limits --- apps/api/app/core/dependencies.py | 25 ++- .../services/auth/api_key_identity_cache.py | 120 +++++++++++++ apps/api/app/services/auth/api_key_service.py | 43 ++++- .../app/services/billing/stripe_service.py | 9 + .../app/services/rate_limit/dependencies.py | 108 +----------- .../app/services/rate_limit/identity_cache.py | 160 +++--------------- 6 files changed, 218 insertions(+), 247 deletions(-) create mode 100644 apps/api/app/services/auth/api_key_identity_cache.py diff --git a/apps/api/app/core/dependencies.py b/apps/api/app/core/dependencies.py index 89d2458f8..2449ee8cf 100644 --- a/apps/api/app/core/dependencies.py +++ b/apps/api/app/core/dependencies.py @@ -4,8 +4,8 @@ from typing import Any import jwt +from app.services.auth.api_key_identity_cache import api_key_identity_cache from app.services.auth.api_key_service import APIKeyService -from app.services.rate_limit.identity_cache import identity_cache from fastapi import Depends, Header, Request from jwt import PyJWKClient from loguru import logger @@ -185,22 +185,20 @@ async def get_current_user_id( # Mode 1: API Key verification (for external clients) if token.startswith("sk_"): + api_key_service = APIKeyService() # Check identity cache first — skip DB on cache hit api_key_hash = hash_api_key(token) try: redis_service = redis_pool_manager.get_redis_service() - cached = await identity_cache.get_cached_identity( - redis_service, - identity_cache._apikey_key(api_key_hash), + cached = await api_key_identity_cache.get_identity( + redis_service, api_key_hash ) if cached is not None: cached_user_id = cached.get("user_id") cached_user_tier = cached.get("user_tier") if cached_user_id and isinstance(cached_user_tier, str): request.state.cached_user_tier = cached_user_tier - request.state.cached_identity_hit = True request.state.user_id = cached_user_id - request.state.api_key_hash = api_key_hash _enforce_guest_api_key_scope(route_path, cached_user_tier) return cached_user_id except PermissionDeniedException: @@ -209,13 +207,22 @@ async def get_current_user_id( pass # Fall through to DB validation # Cache miss — validate via DB - api_key_service = APIKeyService() identity = await api_key_service.validate_api_key_identity(db, token) if identity: request.state.cached_user_tier = identity.user_tier - request.state.cached_identity_hit = False request.state.user_id = identity.user_id - request.state.api_key_hash = identity.key_hash + try: + await api_key_service.cache_api_key_identity( + key_hash=identity.key_hash, + user_id=identity.user_id, + user_tier=identity.user_tier, + expires_at=identity.expires_at, + ) + except Exception: + logger.warning( + "auth: failed to cache API key identity for user_id={}", + identity.user_id, + ) _enforce_guest_api_key_scope(route_path, identity.user_tier) return identity.user_id else: diff --git a/apps/api/app/services/auth/api_key_identity_cache.py b/apps/api/app/services/auth/api_key_identity_cache.py new file mode 100644 index 000000000..a4a0cbae2 --- /dev/null +++ b/apps/api/app/services/auth/api_key_identity_cache.py @@ -0,0 +1,120 @@ +"""Redis-backed API-key authentication identity cache.""" + +import json + +from loguru import logger + +from shared.services.redis.redis_service import RedisService + +_API_KEY_MAX_TTL_SECONDS: int = 3600 + + +class APIKeyIdentityCache: + """Cache validated API-key identities by API-key lookup hash.""" + + @staticmethod + def get_cache_key(api_key_hash: str) -> str: + """Return the Redis key for an API-key hash.""" + return f"identity:apikey:{api_key_hash}" + + @staticmethod + def get_reverse_key(user_id: str) -> str: + """Return the reverse-index Redis key for a user.""" + return f"identity:apikeys:{user_id}" + + async def get_identity( + self, + redis: RedisService, + api_key_hash: str, + ) -> dict[str, str] | None: + """Return cached ``{user_id, user_tier}`` for an API key.""" + try: + raw_identity: object = await redis.get(self.get_cache_key(api_key_hash)) + return self._coerce_identity(raw_identity) + except Exception: + logger.warning( + "api_key_identity_cache: failed to read identity", + ) + return None + + async def set_identity( + self, + redis: RedisService, + api_key_hash: str, + user_id: str, + user_tier: str, + ttl_seconds: int, + ) -> None: + """Cache a validated API-key identity.""" + effective_ttl_seconds: int = min(_API_KEY_MAX_TTL_SECONDS, ttl_seconds) + cache_key: str = self.get_cache_key(api_key_hash) + reverse_key: str = self.get_reverse_key(user_id) + payload: dict[str, str] = {"user_id": user_id, "user_tier": user_tier} + + try: + await redis.set(cache_key, payload, ttl=effective_ttl_seconds) + await redis.sadd(reverse_key, api_key_hash) + current_ttl_seconds = await redis.ttl(reverse_key) + if ( + current_ttl_seconds in (-2, -1) + or current_ttl_seconds < effective_ttl_seconds + ): + await redis.expire(reverse_key, effective_ttl_seconds) + except Exception: + logger.warning( + "api_key_identity_cache: failed to set identity for user_id={}", + user_id, + ) + + async def invalidate_api_key( + self, + redis: RedisService, + user_id: str, + api_key_hash: str, + ) -> None: + """Delete one API-key identity cache entry.""" + try: + await redis.delete(self.get_cache_key(api_key_hash)) + await redis.srem(self.get_reverse_key(user_id), api_key_hash) + except Exception: + logger.warning( + "api_key_identity_cache: failed to invalidate identity for user_id={}", + user_id, + ) + + async def invalidate_user( + self, + redis: RedisService, + user_id: str, + ) -> None: + """Delete all API-key identity cache entries for a user.""" + try: + reverse_key: str = self.get_reverse_key(user_id) + api_key_hashes: set[object] = await redis.smembers(reverse_key) + for api_key_hash in api_key_hashes: + await redis.delete(self.get_cache_key(str(api_key_hash))) + await redis.delete(reverse_key) + except Exception: + logger.warning( + "api_key_identity_cache: failed to invalidate user_id={}", + user_id, + ) + + def _coerce_identity(self, raw_identity: object) -> dict[str, str] | None: + """Return a typed identity payload from a Redis value.""" + parsed_identity = raw_identity + if isinstance(raw_identity, str): + parsed_identity = json.loads(raw_identity) + + if not isinstance(parsed_identity, dict): + return None + + user_id = parsed_identity.get("user_id") + user_tier = parsed_identity.get("user_tier") + if not isinstance(user_id, str) or not isinstance(user_tier, str): + return None + + return {"user_id": user_id, "user_tier": user_tier} + + +api_key_identity_cache = APIKeyIdentityCache() diff --git a/apps/api/app/services/auth/api_key_service.py b/apps/api/app/services/auth/api_key_service.py index 5e81ae592..dd85e4854 100644 --- a/apps/api/app/services/auth/api_key_service.py +++ b/apps/api/app/services/auth/api_key_service.py @@ -2,11 +2,11 @@ import asyncio from dataclasses import dataclass -from datetime import datetime +from datetime import datetime, timezone from typing import List, Optional from app.repositories.api_key_repository import APIKeyRepository -from app.services.rate_limit.identity_cache import identity_cache +from app.services.auth.api_key_identity_cache import api_key_identity_cache from loguru import logger from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession @@ -37,6 +37,7 @@ class APIKeyIdentity: user_id: str user_tier: str key_hash: str + expires_at: datetime | None class APIKeyService: @@ -128,6 +129,24 @@ async def validate_api_key_identity( user_id=user_id, user_tier=user_tier, key_hash=str(api_key_record.key_hash), + expires_at=api_key_record.expires_at, + ) + + async def cache_api_key_identity( + self, + *, + key_hash: str, + user_id: str, + user_tier: str, + expires_at: datetime | None, + ) -> None: + """Cache a validated API-key identity for the auth layer.""" + await api_key_identity_cache.set_identity( + redis_pool_manager.get_redis_service(), + key_hash, + user_id, + user_tier, + ttl_seconds=self._resolve_api_key_cache_ttl_seconds(expires_at), ) async def _resolve_user_tier( @@ -142,6 +161,20 @@ async def _resolve_user_tier( user_tier = result.scalar_one_or_none() return str(user_tier) if user_tier is not None else _DEFAULT_USER_TIER + def _resolve_api_key_cache_ttl_seconds(self, expires_at: datetime | None) -> int: + """Resolve cache TTL for API-key identity without exceeding key expiry.""" + max_ttl_seconds = 3600 + if expires_at is None: + return max_ttl_seconds + + expires_at_utc = expires_at + if expires_at_utc.tzinfo is None: + expires_at_utc = expires_at_utc.replace(tzinfo=timezone.utc) + + now = datetime.now(timezone.utc) + remaining_seconds = int((expires_at_utc - now).total_seconds()) + return max(1, min(max_ttl_seconds, remaining_seconds)) + async def revoke_api_key( self, session: AsyncSession, api_key_id: str, user_id: str ) -> bool: @@ -191,7 +224,7 @@ async def _invalidate_revoked_api_key_cache_best_effort( ) -> None: """Best-effort cache invalidation after a revoke has already been committed.""" try: - await identity_cache.invalidate_apikey( + await api_key_identity_cache.invalidate_api_key( redis_pool_manager.get_redis_service(), user_id, key_hash, @@ -254,7 +287,7 @@ async def regenerate_api_key( await session.commit() # 4. Refresh the cache. - await identity_cache.invalidate_apikey( + await api_key_identity_cache.invalidate_api_key( redis_pool_manager.get_redis_service(), user_id, api_key.key_hash, @@ -328,7 +361,7 @@ async def toggle_api_key( await session.refresh(api_key) if not api_key.is_active: - await identity_cache.invalidate_apikey( + await api_key_identity_cache.invalidate_api_key( redis_pool_manager.get_redis_service(), user_id, api_key.key_hash, diff --git a/apps/api/app/services/billing/stripe_service.py b/apps/api/app/services/billing/stripe_service.py index 3d10b5688..519337e46 100644 --- a/apps/api/app/services/billing/stripe_service.py +++ b/apps/api/app/services/billing/stripe_service.py @@ -5,6 +5,7 @@ import stripe from app.repositories.payment_record_repository import PaymentRecordRepository +from app.services.auth.api_key_identity_cache import api_key_identity_cache from app.services.billing.price_config_service import PriceConfigService from app.services.rate_limit.identity_cache import identity_cache from app.services.rate_limit.tier_service import TierService @@ -389,6 +390,10 @@ async def _handle_checkout_completed( redis_pool_manager.get_redis_service(), user_id, ) + await api_key_identity_cache.invalidate_user( + redis_pool_manager.get_redis_service(), + user_id, + ) logger.info( f"Credits pack purchase succeeded: user_id={user_id}, credits={credits_amount}, price_id={price_id}" @@ -531,6 +536,10 @@ async def _handle_payment_intent_succeeded( redis_pool_manager.get_redis_service(), user_id, ) + await api_key_identity_cache.invalidate_user( + redis_pool_manager.get_redis_service(), + user_id, + ) logger.info( f"buy credits success: user_id={user_id}, credits={credits_amount}, payment_intent_id={payment_intent_id}" diff --git a/apps/api/app/services/rate_limit/dependencies.py b/apps/api/app/services/rate_limit/dependencies.py index d92ecd8fd..eea907b0b 100644 --- a/apps/api/app/services/rate_limit/dependencies.py +++ b/apps/api/app/services/rate_limit/dependencies.py @@ -5,8 +5,9 @@ require_billing_limits -> with_current_user -> get_current_user_id -> get_db -``with_current_user`` resolves identity (user_id + user_tier), caches it in -Redis, and enforces the matched system limit (Layer 0). +``with_current_user`` resolves rate-limit identity (user_id + user_tier), +caches the tier by user_id in Redis, and enforces the matched system limit +(Layer 0). ``require_billing_limits`` enforces billing RPM (Layer 1) when billing is enabled and yields control to the route handler. Concurrency (Layer 2) and @@ -15,7 +16,6 @@ """ import math -from datetime import datetime, timezone from typing import AsyncGenerator from app.core.dependencies import get_current_user_id @@ -40,10 +40,8 @@ ) from shared.core.logging import log_context from shared.core.state_machine.states import JobStatus -from shared.models.database.api_key import APIKey from shared.models.database.job import Job from shared.models.database.user_balance import UserBalance -from shared.utils.api_keys import hash_api_key, is_api_key_token _DEFAULT_TIER: str = "free" _ACTIVE_JOB_STATES: tuple[str, ...] = ( @@ -83,16 +81,6 @@ async def _resolve_user_tier_from_db(user_id: str) -> str: return _DEFAULT_TIER -def _extract_bearer_token(authorization: str | None) -> str | None: - """Extract bearer token from Authorization header.""" - if not authorization: - return None - scheme, _, token = authorization.partition(" ") - if scheme.lower() != "bearer" or not token: - return None - return token - - def _get_route_path(request: Request) -> str: """Return the request path without the application's root_path prefix.""" scope_path: str = request.scope.get("path", request.url.path) @@ -116,30 +104,6 @@ def _get_route_limit_identifier(request: Request) -> str: return _get_route_path(request) -async def _resolve_apikey_cache_ttl_seconds(api_key_hash: str) -> int: - """Resolve cache TTL for API key identity (max 1 hour).""" - max_ttl_seconds = 3600 - try: - async with get_db_context() as db: - result = await db.execute( - select(APIKey.expires_at) - .where(APIKey.key_hash == api_key_hash) - .limit(1) - ) - expires_at = result.scalar_one_or_none() - if expires_at is None: - return max_ttl_seconds - - # APIKey.expires_at is stored as UTC-naive datetime. - if expires_at.tzinfo is None: - expires_at = expires_at.replace(tzinfo=timezone.utc) - now = datetime.now(timezone.utc) - remaining = int((expires_at - now).total_seconds()) - return max(1, min(max_ttl_seconds, remaining)) - except Exception: - return max_ttl_seconds - - # --------------------------------------------------------------------------- # with_current_user -- Layer 0 (matched system limit) # --------------------------------------------------------------------------- @@ -165,72 +129,18 @@ async def with_current_user( # -- Resolve user_tier (cache -> DB fallback) -- user_tier: str | None = getattr(request.state, "cached_user_tier", None) stashed_user_id: str | None = getattr(request.state, "user_id", None) - cached_identity_hit: bool | None = getattr( - request.state, "cached_identity_hit", None - ) - if isinstance(user_tier, str) and stashed_user_id == user_id: - if cached_identity_hit is False: - token = _extract_bearer_token(request.headers.get("authorization")) - api_key_hash = getattr(request.state, "api_key_hash", None) - if not isinstance(api_key_hash, str) and is_api_key_token(token): - api_key_hash = hash_api_key(token) - if isinstance(api_key_hash, str): - try: - ttl_seconds = await _resolve_apikey_cache_ttl_seconds(api_key_hash) - await identity_cache.set_apikey_identity( - redis_service, - api_key_hash, - user_id, - user_tier, - ttl_seconds=ttl_seconds, - ) - except Exception: - logger.warning( - "rate_limit: failed to backfill API key identity cache " - "for user_id={}", - user_id, - ) - else: - token = _extract_bearer_token(request.headers.get("authorization")) - api_key_hash: str | None = None - if is_api_key_token(token): - api_key_hash = hash_api_key(token) - cache_key: str = ( - identity_cache._apikey_key(api_key_hash) - if isinstance(api_key_hash, str) - else identity_cache._jwt_key(user_id) - ) + if not isinstance(user_tier, str) or stashed_user_id != user_id: try: - if isinstance(api_key_hash, str): - cached = await identity_cache.get_cached_identity( - redis_service, - identity_cache._apikey_key(api_key_hash), - ) - else: - cached = await identity_cache.get_cached_identity( - redis_service, - cache_key, - ) + cached = await identity_cache.get_cached_identity( + redis_service, + identity_cache.get_user_key(user_id), + ) if cached is not None: user_tier = cached.get("user_tier", _DEFAULT_TIER) - if isinstance(api_key_hash, str): - request.state.api_key_hash = api_key_hash else: user_tier = await _resolve_user_tier_from_db(user_id) - if isinstance(api_key_hash, str): - ttl_seconds = await _resolve_apikey_cache_ttl_seconds(api_key_hash) - await identity_cache.set_apikey_identity( - redis_service, - api_key_hash, - user_id, - user_tier, - ttl_seconds=ttl_seconds, - ) - else: - await identity_cache.set_jwt_identity( - redis_service, user_id, user_tier - ) + await identity_cache.set_jwt_identity(redis_service, user_id, user_tier) except Exception: logger.warning( "rate_limit: Redis error during identity resolution, " diff --git a/apps/api/app/services/rate_limit/identity_cache.py b/apps/api/app/services/rate_limit/identity_cache.py index 9e6ffae01..4dab693bd 100644 --- a/apps/api/app/services/rate_limit/identity_cache.py +++ b/apps/api/app/services/rate_limit/identity_cache.py @@ -1,69 +1,30 @@ -""" -Redis-backed identity cache for user_id + user_tier resolution. - -Caches the mapping from authentication credentials (JWT user_id or API key hash) -to the resolved identity (user_id, user_tier) so that -tier lookups do not hit the database on every request. - -Key patterns (all prefixed with REDIS_KEY_PREFIX from config): - JWT: {REDIS_KEY_PREFIX}identity:jwt:{user_id} - API key: {REDIS_KEY_PREFIX}identity:apikey:{api_key_hash} - Reverse: {REDIS_KEY_PREFIX}identity:apikeys:{user_id} -""" +"""Redis-backed rate-limit identity cache for user_id + user_tier.""" import json -from typing import Optional -from app.services.rate_limit.config import REDIS_KEY_PREFIX from loguru import logger from shared.services.redis.redis_service import RedisService -# Default TTL for JWT identity cache entries (1 hour). _JWT_TTL_SECONDS: int = 3600 -# Upper bound TTL for API-key identity cache entries (1 hour). -_APIKEY_MAX_TTL_SECONDS: int = 3600 - class IdentityCache: - """Redis-backed identity cache for user_id + user_tier resolution.""" - - # ------------------------------------------------------------------ - # Key builders - # ------------------------------------------------------------------ - - @staticmethod - def _jwt_key(user_id: str) -> str: - return f"{REDIS_KEY_PREFIX}identity:jwt:{user_id}" - - @staticmethod - def _apikey_key(api_key_hash: str) -> str: - return f"{REDIS_KEY_PREFIX}identity:apikey:{api_key_hash}" + """Cache resolved rate-limit identity by user_id.""" @staticmethod - def _reverse_key(user_id: str) -> str: - return f"{REDIS_KEY_PREFIX}identity:apikeys:{user_id}" - - # ------------------------------------------------------------------ - # Read - # ------------------------------------------------------------------ + def get_user_key(user_id: str) -> str: + return f"identity:user:{user_id}" async def get_cached_identity( self, redis: RedisService, cache_key: str, - ) -> Optional[dict]: + ) -> dict[str, str] | None: """Return cached ``{user_id, user_tier}`` or ``None`` on miss.""" try: - raw: Optional[str] = await redis.get(cache_key) - if raw is None: - return None - # RedisService.get already attempts JSON parse, but the - # value may come back as a dict directly. - if isinstance(raw, dict): - return raw - return json.loads(raw) + raw_identity: object = await redis.get(cache_key) + return self._coerce_identity(raw_identity) except Exception: logger.warning( "identity_cache: failed to read cache_key={}", @@ -71,19 +32,15 @@ async def get_cached_identity( ) return None - # ------------------------------------------------------------------ - # Write -- JWT - # ------------------------------------------------------------------ - async def set_jwt_identity( self, redis: RedisService, user_id: str, user_tier: str, ) -> None: - """Cache identity for a JWT-authenticated user (1 hr TTL).""" - key: str = self._jwt_key(user_id) - payload: dict = {"user_id": user_id, "user_tier": user_tier} + """Cache rate-limit identity for a user.""" + key: str = self.get_user_key(user_id) + payload: dict[str, str] = {"user_id": user_id, "user_tier": user_tier} try: await redis.set(key, payload, ttl=_JWT_TTL_SECONDS) except Exception: @@ -92,100 +49,35 @@ async def set_jwt_identity( user_id, ) - # ------------------------------------------------------------------ - # Write -- API key - # ------------------------------------------------------------------ - - async def set_apikey_identity( - self, - redis: RedisService, - api_key_hash: str, - user_id: str, - user_tier: str, - ttl_seconds: int, - ) -> None: - """Cache identity for an API-key-authenticated user. - - TTL is ``min(APIKEY_MAX_TTL, api_key_remaining_ttl)`` so the - cache never outlives the key itself. Also maintains a reverse - index (SET) of all cached API-key hashes per user for bulk - invalidation. - """ - effective_ttl: int = min(_APIKEY_MAX_TTL_SECONDS, ttl_seconds) - key: str = self._apikey_key(api_key_hash) - payload: dict = {"user_id": user_id, "user_tier": user_tier} - try: - await redis.set(key, payload, ttl=effective_ttl) - # Maintain reverse index so invalidate_user can find all - # API-key cache entries belonging to this user. - reverse_key: str = self._reverse_key(user_id) - await redis.sadd(reverse_key, api_key_hash) - # Keep reverse index TTL at least as long as the longest - # surviving API-key cache entry for this user. - current_ttl = await redis.ttl(reverse_key) - if current_ttl in (-2, -1) or current_ttl < effective_ttl: - await redis.expire(reverse_key, effective_ttl) - except Exception: - logger.warning( - "identity_cache: failed to set apikey identity " - "api_key_hash={}, user_id={}", - api_key_hash, - user_id, - ) - - # ------------------------------------------------------------------ - # Invalidation - # ------------------------------------------------------------------ - async def invalidate_user( self, redis: RedisService, user_id: str, ) -> None: - """Full invalidation: JWT cache + all API-key caches + reverse index.""" + """Delete cached rate-limit identity for a user.""" try: - # 1. Delete JWT cache - jwt_key: str = self._jwt_key(user_id) - await redis.delete(jwt_key) - - # 2. Collect all cached API-key hashes from reverse index - reverse_key: str = self._reverse_key(user_id) - api_key_hashes: set = await redis.smembers(reverse_key) - - # 3. Delete each API-key cache entry - for api_key_hash in api_key_hashes: - apikey_key: str = self._apikey_key(str(api_key_hash)) - await redis.delete(apikey_key) - - # 4. Delete the reverse index itself - await redis.delete(reverse_key) + await redis.delete(self.get_user_key(user_id)) except Exception: logger.warning( "identity_cache: failed to invalidate user_id={}", user_id, ) - async def invalidate_apikey( - self, - redis: RedisService, - user_id: str, - api_key_hash: str, - ) -> None: - """Delete a single API-key cache entry and remove from reverse index.""" - try: - apikey_key: str = self._apikey_key(api_key_hash) - await redis.delete(apikey_key) + def _coerce_identity(self, raw_identity: object) -> dict[str, str] | None: + """Return a typed identity payload from a Redis value.""" + parsed_identity = raw_identity + if isinstance(raw_identity, str): + parsed_identity = json.loads(raw_identity) - reverse_key: str = self._reverse_key(user_id) - await redis.srem(reverse_key, api_key_hash) - except Exception: - logger.warning( - "identity_cache: failed to invalidate apikey " - "api_key_hash={}, user_id={}", - api_key_hash, - user_id, - ) + if not isinstance(parsed_identity, dict): + return None + + user_id = parsed_identity.get("user_id") + user_tier = parsed_identity.get("user_tier") + if not isinstance(user_id, str) or not isinstance(user_tier, str): + return None + + return {"user_id": user_id, "user_tier": user_tier} -# Module-level singleton so callers can import directly. identity_cache = IdentityCache() From 132bd2bb0b7251cc6e48f2f062f726a65427c3c8 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 03:34:19 +0000 Subject: [PATCH 12/32] refactor: centralize api key identity lookup --- apps/api/app/core/dependencies.py | 39 +--------- apps/api/app/services/auth/api_key_service.py | 72 ++++++++++++++----- 2 files changed, 58 insertions(+), 53 deletions(-) diff --git a/apps/api/app/core/dependencies.py b/apps/api/app/core/dependencies.py index 2449ee8cf..69f57ea2a 100644 --- a/apps/api/app/core/dependencies.py +++ b/apps/api/app/core/dependencies.py @@ -4,20 +4,18 @@ from typing import Any import jwt -from app.services.auth.api_key_identity_cache import api_key_identity_cache from app.services.auth.api_key_service import APIKeyService from fastapi import Depends, Header, Request from jwt import PyJWKClient from loguru import logger from sqlalchemy.ext.asyncio import AsyncSession -from shared.core.config import redis_pool_manager, settings +from shared.core.config import settings from shared.core.database import get_db from shared.core.exceptions.domain_exceptions import ( AuthException, PermissionDeniedException, ) -from shared.utils.api_keys import hash_api_key # Standard JWKS endpoint path (fixed, following OpenID Connect convention) JWKS_ENDPOINT_PATH = "/api/auth/jwks" @@ -186,43 +184,10 @@ async def get_current_user_id( # Mode 1: API Key verification (for external clients) if token.startswith("sk_"): api_key_service = APIKeyService() - # Check identity cache first — skip DB on cache hit - api_key_hash = hash_api_key(token) - try: - redis_service = redis_pool_manager.get_redis_service() - cached = await api_key_identity_cache.get_identity( - redis_service, api_key_hash - ) - if cached is not None: - cached_user_id = cached.get("user_id") - cached_user_tier = cached.get("user_tier") - if cached_user_id and isinstance(cached_user_tier, str): - request.state.cached_user_tier = cached_user_tier - request.state.user_id = cached_user_id - _enforce_guest_api_key_scope(route_path, cached_user_tier) - return cached_user_id - except PermissionDeniedException: - raise - except Exception: - pass # Fall through to DB validation - - # Cache miss — validate via DB - identity = await api_key_service.validate_api_key_identity(db, token) + identity = await api_key_service.get_identity(db, token) if identity: request.state.cached_user_tier = identity.user_tier request.state.user_id = identity.user_id - try: - await api_key_service.cache_api_key_identity( - key_hash=identity.key_hash, - user_id=identity.user_id, - user_tier=identity.user_tier, - expires_at=identity.expires_at, - ) - except Exception: - logger.warning( - "auth: failed to cache API key identity for user_id={}", - identity.user_id, - ) _enforce_guest_api_key_scope(route_path, identity.user_tier) return identity.user_id else: diff --git a/apps/api/app/services/auth/api_key_service.py b/apps/api/app/services/auth/api_key_service.py index dd85e4854..7752d8d8c 100644 --- a/apps/api/app/services/auth/api_key_service.py +++ b/apps/api/app/services/auth/api_key_service.py @@ -36,7 +36,6 @@ class APIKeyIdentity: user_id: str user_tier: str - key_hash: str expires_at: datetime | None @@ -106,16 +105,52 @@ async def validate_api_key( self, session: AsyncSession, api_key: str ) -> Optional[str]: """Validate API key against DB, return user_id or None.""" - identity = await self.validate_api_key_identity(session, api_key) + identity = await self.get_identity(session, api_key) return identity.user_id if identity is not None else None - async def validate_api_key_identity( + async def get_identity( self, session: AsyncSession, api_key: str, ) -> Optional[APIKeyIdentity]: - """Validate API key and return the authenticated identity.""" + """Return API-key identity, using Redis cache before DB fallback.""" key_hash = hash_api_key(api_key) + cached_identity = await self._get_cached_identity(key_hash) + if cached_identity is not None: + return cached_identity + + identity = await self._get_database_identity(session, api_key, key_hash) + if identity is None: + return None + + await self._cache_api_key_identity( + key_hash=key_hash, + identity=identity, + ) + return identity + + async def _get_cached_identity(self, key_hash: str) -> APIKeyIdentity | None: + """Return cached API-key identity, or None on miss/cache failure.""" + cached_identity = await api_key_identity_cache.get_identity( + redis_pool_manager.get_redis_service(), + key_hash, + ) + if cached_identity is None: + return None + + return APIKeyIdentity( + user_id=cached_identity["user_id"], + user_tier=cached_identity["user_tier"], + expires_at=None, + ) + + async def _get_database_identity( + self, + session: AsyncSession, + api_key: str, + key_hash: str, + ) -> APIKeyIdentity | None: + """Validate API key against the database.""" api_key_record = await self.repository.get_by_key_hash(session, key_hash) if not api_key_record or not api_key_record.is_valid(): @@ -128,26 +163,31 @@ async def validate_api_key_identity( return APIKeyIdentity( user_id=user_id, user_tier=user_tier, - key_hash=str(api_key_record.key_hash), expires_at=api_key_record.expires_at, ) - async def cache_api_key_identity( + async def _cache_api_key_identity( self, *, key_hash: str, - user_id: str, - user_tier: str, - expires_at: datetime | None, + identity: APIKeyIdentity, ) -> None: """Cache a validated API-key identity for the auth layer.""" - await api_key_identity_cache.set_identity( - redis_pool_manager.get_redis_service(), - key_hash, - user_id, - user_tier, - ttl_seconds=self._resolve_api_key_cache_ttl_seconds(expires_at), - ) + try: + await api_key_identity_cache.set_identity( + redis_pool_manager.get_redis_service(), + key_hash, + identity.user_id, + identity.user_tier, + ttl_seconds=self._resolve_api_key_cache_ttl_seconds( + identity.expires_at + ), + ) + except Exception: + logger.warning( + "api_key_service: failed to cache identity for user_id={}", + identity.user_id, + ) async def _resolve_user_tier( self, From 4d821d9dcb1e4b3194dce80c11e3b24f5ff229da Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 03:42:55 +0000 Subject: [PATCH 13/32] refactor: invalidate tier identity caches in tier service --- .../app/services/billing/stripe_service.py | 22 +------------------ .../app/services/rate_limit/tier_service.py | 17 ++++++++++++++ 2 files changed, 18 insertions(+), 21 deletions(-) diff --git a/apps/api/app/services/billing/stripe_service.py b/apps/api/app/services/billing/stripe_service.py index 519337e46..be1792a2e 100644 --- a/apps/api/app/services/billing/stripe_service.py +++ b/apps/api/app/services/billing/stripe_service.py @@ -5,14 +5,12 @@ import stripe from app.repositories.payment_record_repository import PaymentRecordRepository -from app.services.auth.api_key_identity_cache import api_key_identity_cache from app.services.billing.price_config_service import PriceConfigService -from app.services.rate_limit.identity_cache import identity_cache from app.services.rate_limit.tier_service import TierService from sqlalchemy import func, select from sqlalchemy.ext.asyncio import AsyncSession -from shared.core.config import redis_pool_manager, settings +from shared.core.config import settings from shared.core.exceptions.domain_exceptions import ( AuthException, KnowhereException, @@ -386,15 +384,6 @@ async def _handle_checkout_completed( await db.commit() await db.refresh(payment_record) - await identity_cache.invalidate_user( - redis_pool_manager.get_redis_service(), - user_id, - ) - await api_key_identity_cache.invalidate_user( - redis_pool_manager.get_redis_service(), - user_id, - ) - logger.info( f"Credits pack purchase succeeded: user_id={user_id}, credits={credits_amount}, price_id={price_id}" ) @@ -532,15 +521,6 @@ async def _handle_payment_intent_succeeded( await db.commit() await db.refresh(payment_record) - await identity_cache.invalidate_user( - redis_pool_manager.get_redis_service(), - user_id, - ) - await api_key_identity_cache.invalidate_user( - redis_pool_manager.get_redis_service(), - user_id, - ) - logger.info( f"buy credits success: user_id={user_id}, credits={credits_amount}, payment_intent_id={payment_intent_id}" ) diff --git a/apps/api/app/services/rate_limit/tier_service.py b/apps/api/app/services/rate_limit/tier_service.py index a9d73a13e..ef3643481 100644 --- a/apps/api/app/services/rate_limit/tier_service.py +++ b/apps/api/app/services/rate_limit/tier_service.py @@ -5,12 +5,15 @@ from typing import Optional +from app.services.auth.api_key_identity_cache import api_key_identity_cache from app.services.rate_limit.config import RateLimitConfig from app.services.rate_limit.data_structures import TierLimits +from app.services.rate_limit.identity_cache import identity_cache from loguru import logger from sqlalchemy import func, select, update from sqlalchemy.ext.asyncio import AsyncSession +from shared.core.config import redis_pool_manager from shared.models.database.payment_record import PaymentRecord from shared.models.database.tier_limit import TierLimit from shared.models.database.user_balance import UserBalance @@ -61,6 +64,7 @@ async def refresh_tier(user_id: str, session: AsyncSession) -> str: .values(user_tier=new_tier) ) await session.execute(stmt_update) + await TierService._invalidate_identity_caches(user_id) logger.info( "Tier refreshed: user_id=%s total_micro=%d new_tier=%s", @@ -78,3 +82,16 @@ def get_limits(user_tier: str) -> Optional[TierLimits]: """ config = RateLimitConfig.get_instance() return config.tier_map.get(user_tier) + + @staticmethod + async def _invalidate_identity_caches(user_id: str) -> None: + """Invalidate identity caches that include user_tier.""" + try: + redis_service = redis_pool_manager.get_redis_service() + await identity_cache.invalidate_user(redis_service, user_id) + await api_key_identity_cache.invalidate_user(redis_service, user_id) + except Exception: + logger.warning( + "Tier refresh cache invalidation failed for user_id={}", + user_id, + ) From 6573836454c4cdc873811896aa2135fd6f29d7b2 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 03:49:47 +0000 Subject: [PATCH 14/32] refactor: keep api key cache auth only --- .../services/auth/api_key_identity_cache.py | 60 ++++++++++--------- apps/api/app/services/auth/api_key_service.py | 35 +++++++---- .../app/services/rate_limit/tier_service.py | 2 - 3 files changed, 56 insertions(+), 41 deletions(-) diff --git a/apps/api/app/services/auth/api_key_identity_cache.py b/apps/api/app/services/auth/api_key_identity_cache.py index a4a0cbae2..a985cdb31 100644 --- a/apps/api/app/services/auth/api_key_identity_cache.py +++ b/apps/api/app/services/auth/api_key_identity_cache.py @@ -1,4 +1,4 @@ -"""Redis-backed API-key authentication identity cache.""" +"""Redis-backed API-key authentication user cache.""" import json @@ -10,7 +10,7 @@ class APIKeyIdentityCache: - """Cache validated API-key identities by API-key lookup hash.""" + """Cache validated API-key user IDs by API-key lookup hash.""" @staticmethod def get_cache_key(api_key_hash: str) -> str: @@ -22,37 +22,35 @@ def get_reverse_key(user_id: str) -> str: """Return the reverse-index Redis key for a user.""" return f"identity:apikeys:{user_id}" - async def get_identity( + async def get_user_id( self, redis: RedisService, api_key_hash: str, - ) -> dict[str, str] | None: - """Return cached ``{user_id, user_tier}`` for an API key.""" + ) -> str | None: + """Return cached user_id for an API key.""" try: - raw_identity: object = await redis.get(self.get_cache_key(api_key_hash)) - return self._coerce_identity(raw_identity) + raw_user_id: object = await redis.get(self.get_cache_key(api_key_hash)) + return self._coerce_user_id(raw_user_id) except Exception: logger.warning( - "api_key_identity_cache: failed to read identity", + "api_key_identity_cache: failed to read user", ) return None - async def set_identity( + async def set_user_id( self, redis: RedisService, api_key_hash: str, user_id: str, - user_tier: str, ttl_seconds: int, ) -> None: - """Cache a validated API-key identity.""" + """Cache a validated API-key user ID.""" effective_ttl_seconds: int = min(_API_KEY_MAX_TTL_SECONDS, ttl_seconds) cache_key: str = self.get_cache_key(api_key_hash) reverse_key: str = self.get_reverse_key(user_id) - payload: dict[str, str] = {"user_id": user_id, "user_tier": user_tier} try: - await redis.set(cache_key, payload, ttl=effective_ttl_seconds) + await redis.set(cache_key, user_id, ttl=effective_ttl_seconds) await redis.sadd(reverse_key, api_key_hash) current_ttl_seconds = await redis.ttl(reverse_key) if ( @@ -62,7 +60,7 @@ async def set_identity( await redis.expire(reverse_key, effective_ttl_seconds) except Exception: logger.warning( - "api_key_identity_cache: failed to set identity for user_id={}", + "api_key_identity_cache: failed to set user for user_id={}", user_id, ) @@ -100,21 +98,25 @@ async def invalidate_user( user_id, ) - def _coerce_identity(self, raw_identity: object) -> dict[str, str] | None: - """Return a typed identity payload from a Redis value.""" - parsed_identity = raw_identity - if isinstance(raw_identity, str): - parsed_identity = json.loads(raw_identity) - - if not isinstance(parsed_identity, dict): - return None - - user_id = parsed_identity.get("user_id") - user_tier = parsed_identity.get("user_tier") - if not isinstance(user_id, str) or not isinstance(user_tier, str): - return None - - return {"user_id": user_id, "user_tier": user_tier} + def _coerce_user_id(self, raw_user_id: object) -> str | None: + """Return a typed user ID from current or legacy Redis values.""" + if isinstance(raw_user_id, str): + try: + parsed_user_id = json.loads(raw_user_id) + except json.JSONDecodeError: + return raw_user_id + else: + parsed_user_id = raw_user_id + + if isinstance(parsed_user_id, str): + return parsed_user_id + + if isinstance(parsed_user_id, dict): + legacy_user_id = parsed_user_id.get("user_id") + if isinstance(legacy_user_id, str): + return legacy_user_id + + return None api_key_identity_cache = APIKeyIdentityCache() diff --git a/apps/api/app/services/auth/api_key_service.py b/apps/api/app/services/auth/api_key_service.py index 7752d8d8c..b3c5c2cd2 100644 --- a/apps/api/app/services/auth/api_key_service.py +++ b/apps/api/app/services/auth/api_key_service.py @@ -7,6 +7,7 @@ from app.repositories.api_key_repository import APIKeyRepository from app.services.auth.api_key_identity_cache import api_key_identity_cache +from app.services.rate_limit.identity_cache import identity_cache from loguru import logger from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession @@ -130,17 +131,17 @@ async def get_identity( return identity async def _get_cached_identity(self, key_hash: str) -> APIKeyIdentity | None: - """Return cached API-key identity, or None on miss/cache failure.""" - cached_identity = await api_key_identity_cache.get_identity( + """Return cached API-key user identity, or None on miss/cache failure.""" + user_id = await api_key_identity_cache.get_user_id( redis_pool_manager.get_redis_service(), key_hash, ) - if cached_identity is None: + if user_id is None: return None return APIKeyIdentity( - user_id=cached_identity["user_id"], - user_tier=cached_identity["user_tier"], + user_id=user_id, + user_tier=await self._get_user_tier(user_id), expires_at=None, ) @@ -158,7 +159,7 @@ async def _get_database_identity( self._schedule_last_used_update(str(api_key_record.id)) user_id = str(api_key_record.user_id) - user_tier = await self._resolve_user_tier(session, user_id) + user_tier = await self._resolve_user_tier_from_db(session, user_id) return APIKeyIdentity( user_id=user_id, @@ -172,13 +173,12 @@ async def _cache_api_key_identity( key_hash: str, identity: APIKeyIdentity, ) -> None: - """Cache a validated API-key identity for the auth layer.""" + """Cache the validated API-key user ID for the auth layer.""" try: - await api_key_identity_cache.set_identity( + await api_key_identity_cache.set_user_id( redis_pool_manager.get_redis_service(), key_hash, identity.user_id, - identity.user_tier, ttl_seconds=self._resolve_api_key_cache_ttl_seconds( identity.expires_at ), @@ -189,7 +189,22 @@ async def _cache_api_key_identity( identity.user_id, ) - async def _resolve_user_tier( + async def _get_user_tier(self, user_id: str) -> str: + """Return user tier from rate-limit identity cache or DB fallback.""" + redis_service = redis_pool_manager.get_redis_service() + cached_identity = await identity_cache.get_cached_identity( + redis_service, + identity_cache.get_user_key(user_id), + ) + if cached_identity is not None: + return cached_identity["user_tier"] + + async with get_db_context() as session: + user_tier = await self._resolve_user_tier_from_db(session, user_id) + await identity_cache.set_jwt_identity(redis_service, user_id, user_tier) + return user_tier + + async def _resolve_user_tier_from_db( self, session: AsyncSession, user_id: str, diff --git a/apps/api/app/services/rate_limit/tier_service.py b/apps/api/app/services/rate_limit/tier_service.py index ef3643481..e86a4f782 100644 --- a/apps/api/app/services/rate_limit/tier_service.py +++ b/apps/api/app/services/rate_limit/tier_service.py @@ -5,7 +5,6 @@ from typing import Optional -from app.services.auth.api_key_identity_cache import api_key_identity_cache from app.services.rate_limit.config import RateLimitConfig from app.services.rate_limit.data_structures import TierLimits from app.services.rate_limit.identity_cache import identity_cache @@ -89,7 +88,6 @@ async def _invalidate_identity_caches(user_id: str) -> None: try: redis_service = redis_pool_manager.get_redis_service() await identity_cache.invalidate_user(redis_service, user_id) - await api_key_identity_cache.invalidate_user(redis_service, user_id) except Exception: logger.warning( "Tier refresh cache invalidation failed for user_id={}", From 9446b797361755613cf51d575a1722b442cd8dd1 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 03:56:37 +0000 Subject: [PATCH 15/32] refactor: rename user tier cache methods --- apps/api/app/services/auth/api_key_service.py | 7 ++----- .../app/services/rate_limit/dependencies.py | 7 ++----- .../app/services/rate_limit/identity_cache.py | 21 ++++++++++--------- 3 files changed, 15 insertions(+), 20 deletions(-) diff --git a/apps/api/app/services/auth/api_key_service.py b/apps/api/app/services/auth/api_key_service.py index b3c5c2cd2..fc25bc12e 100644 --- a/apps/api/app/services/auth/api_key_service.py +++ b/apps/api/app/services/auth/api_key_service.py @@ -192,16 +192,13 @@ async def _cache_api_key_identity( async def _get_user_tier(self, user_id: str) -> str: """Return user tier from rate-limit identity cache or DB fallback.""" redis_service = redis_pool_manager.get_redis_service() - cached_identity = await identity_cache.get_cached_identity( - redis_service, - identity_cache.get_user_key(user_id), - ) + cached_identity = await identity_cache.get_user_tier(redis_service, user_id) if cached_identity is not None: return cached_identity["user_tier"] async with get_db_context() as session: user_tier = await self._resolve_user_tier_from_db(session, user_id) - await identity_cache.set_jwt_identity(redis_service, user_id, user_tier) + await identity_cache.set_user_tier(redis_service, user_id, user_tier) return user_tier async def _resolve_user_tier_from_db( diff --git a/apps/api/app/services/rate_limit/dependencies.py b/apps/api/app/services/rate_limit/dependencies.py index eea907b0b..6824d66dc 100644 --- a/apps/api/app/services/rate_limit/dependencies.py +++ b/apps/api/app/services/rate_limit/dependencies.py @@ -132,15 +132,12 @@ async def with_current_user( if not isinstance(user_tier, str) or stashed_user_id != user_id: try: - cached = await identity_cache.get_cached_identity( - redis_service, - identity_cache.get_user_key(user_id), - ) + cached = await identity_cache.get_user_tier(redis_service, user_id) if cached is not None: user_tier = cached.get("user_tier", _DEFAULT_TIER) else: user_tier = await _resolve_user_tier_from_db(user_id) - await identity_cache.set_jwt_identity(redis_service, user_id, user_tier) + await identity_cache.set_user_tier(redis_service, user_id, user_tier) except Exception: logger.warning( "rate_limit: Redis error during identity resolution, " diff --git a/apps/api/app/services/rate_limit/identity_cache.py b/apps/api/app/services/rate_limit/identity_cache.py index 4dab693bd..5c67b0acc 100644 --- a/apps/api/app/services/rate_limit/identity_cache.py +++ b/apps/api/app/services/rate_limit/identity_cache.py @@ -6,22 +6,23 @@ from shared.services.redis.redis_service import RedisService -_JWT_TTL_SECONDS: int = 3600 +_USER_TIER_TTL_SECONDS: int = 3600 class IdentityCache: - """Cache resolved rate-limit identity by user_id.""" + """Cache resolved rate-limit user tier by user_id.""" @staticmethod - def get_user_key(user_id: str) -> str: + def get_user_tier_key(user_id: str) -> str: return f"identity:user:{user_id}" - async def get_cached_identity( + async def get_user_tier( self, redis: RedisService, - cache_key: str, + user_id: str, ) -> dict[str, str] | None: """Return cached ``{user_id, user_tier}`` or ``None`` on miss.""" + cache_key = self.get_user_tier_key(user_id) try: raw_identity: object = await redis.get(cache_key) return self._coerce_identity(raw_identity) @@ -32,20 +33,20 @@ async def get_cached_identity( ) return None - async def set_jwt_identity( + async def set_user_tier( self, redis: RedisService, user_id: str, user_tier: str, ) -> None: """Cache rate-limit identity for a user.""" - key: str = self.get_user_key(user_id) + key: str = self.get_user_tier_key(user_id) payload: dict[str, str] = {"user_id": user_id, "user_tier": user_tier} try: - await redis.set(key, payload, ttl=_JWT_TTL_SECONDS) + await redis.set(key, payload, ttl=_USER_TIER_TTL_SECONDS) except Exception: logger.warning( - "identity_cache: failed to set jwt identity user_id={}", + "identity_cache: failed to set user tier for user_id={}", user_id, ) @@ -56,7 +57,7 @@ async def invalidate_user( ) -> None: """Delete cached rate-limit identity for a user.""" try: - await redis.delete(self.get_user_key(user_id)) + await redis.delete(self.get_user_tier_key(user_id)) except Exception: logger.warning( "identity_cache: failed to invalidate user_id={}", From f814d82eb6d679643c30e5208e9998c11f7769e7 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 04:13:24 +0000 Subject: [PATCH 16/32] revert: restore api key and local bootstrap changes --- README.md | 4 - apps/api/app/core/dependencies.py | 29 ++- .../services/auth/api_key_identity_cache.py | 122 ------------ apps/api/app/services/auth/api_key_service.py | 138 +++----------- .../app/services/billing/stripe_service.py | 13 +- .../guest/guest_registration_service.py | 14 +- .../app/services/rate_limit/dependencies.py | 98 +++++++++- .../app/services/rate_limit/identity_cache.py | 173 ++++++++++++++---- .../app/services/rate_limit/tier_service.py | 15 -- apps/api/scripts/bootstrap_local_dev.py | 2 +- apps/api/scripts/init_user.py | 45 ++--- .../scripts/local_dev_bootstrap_service.py | 25 ++- .../tests/contract/test_api_key_contract.py | 34 ---- apps/api/tests/support/contract_database.py | 6 +- deploy/local-dev/README.md | 2 +- .../shared/testing/contract_runtime.py | 2 +- .../shared/tests/utils/test_api_keys.py | 34 ---- .../shared-python/shared/utils/api_keys.py | 30 --- 18 files changed, 340 insertions(+), 446 deletions(-) delete mode 100644 apps/api/app/services/auth/api_key_identity_cache.py delete mode 100644 packages/shared-python/shared/tests/utils/test_api_keys.py delete mode 100644 packages/shared-python/shared/utils/api_keys.py diff --git a/README.md b/README.md index 3057d4143..16e463de2 100644 --- a/README.md +++ b/README.md @@ -83,10 +83,6 @@ uv run --python 3.11 python -m alembic upgrade heads uv run --python 3.11 python scripts/init_user.py --email you@example.com ``` -Pass `--api-key-output-file ./standalone-api-key.txt` if you need the generated -plaintext key written to a local `0600` file. The default console output only -reports that the credential was created. - If you plan to use the dashboard, start the combined self-hosted stack and register through the dashboard instead of using `scripts/init_user.py`. diff --git a/apps/api/app/core/dependencies.py b/apps/api/app/core/dependencies.py index 69f57ea2a..ac349e34f 100644 --- a/apps/api/app/core/dependencies.py +++ b/apps/api/app/core/dependencies.py @@ -1,3 +1,4 @@ +import hashlib import threading from datetime import timedelta from fnmatch import fnmatch @@ -5,12 +6,13 @@ import jwt from app.services.auth.api_key_service import APIKeyService +from app.services.rate_limit.identity_cache import identity_cache from fastapi import Depends, Header, Request from jwt import PyJWKClient from loguru import logger from sqlalchemy.ext.asyncio import AsyncSession -from shared.core.config import settings +from shared.core.config import redis_pool_manager, settings from shared.core.database import get_db from shared.core.exceptions.domain_exceptions import ( AuthException, @@ -183,10 +185,33 @@ async def get_current_user_id( # Mode 1: API Key verification (for external clients) if token.startswith("sk_"): + # Check identity cache first — skip DB on cache hit + api_key_hash = hashlib.sha256(token.encode()).hexdigest() + try: + cached = await identity_cache.get_cached_identity( + redis_pool_manager.get_redis_service(), + identity_cache._apikey_key(api_key_hash), + ) + if cached is not None: + cached_user_id = cached.get("user_id") + cached_user_tier = cached.get("user_tier") + if cached_user_id and isinstance(cached_user_tier, str): + request.state.cached_user_tier = cached_user_tier + request.state.cached_identity_hit = True + request.state.user_id = cached_user_id + _enforce_guest_api_key_scope(route_path, cached_user_tier) + return cached_user_id + except PermissionDeniedException: + raise + except Exception: + pass # Fall through to DB validation + + # Cache miss — validate via DB api_key_service = APIKeyService() - identity = await api_key_service.get_identity(db, token) + identity = await api_key_service.validate_api_key_identity(db, token) if identity: request.state.cached_user_tier = identity.user_tier + request.state.cached_identity_hit = False request.state.user_id = identity.user_id _enforce_guest_api_key_scope(route_path, identity.user_tier) return identity.user_id diff --git a/apps/api/app/services/auth/api_key_identity_cache.py b/apps/api/app/services/auth/api_key_identity_cache.py deleted file mode 100644 index a985cdb31..000000000 --- a/apps/api/app/services/auth/api_key_identity_cache.py +++ /dev/null @@ -1,122 +0,0 @@ -"""Redis-backed API-key authentication user cache.""" - -import json - -from loguru import logger - -from shared.services.redis.redis_service import RedisService - -_API_KEY_MAX_TTL_SECONDS: int = 3600 - - -class APIKeyIdentityCache: - """Cache validated API-key user IDs by API-key lookup hash.""" - - @staticmethod - def get_cache_key(api_key_hash: str) -> str: - """Return the Redis key for an API-key hash.""" - return f"identity:apikey:{api_key_hash}" - - @staticmethod - def get_reverse_key(user_id: str) -> str: - """Return the reverse-index Redis key for a user.""" - return f"identity:apikeys:{user_id}" - - async def get_user_id( - self, - redis: RedisService, - api_key_hash: str, - ) -> str | None: - """Return cached user_id for an API key.""" - try: - raw_user_id: object = await redis.get(self.get_cache_key(api_key_hash)) - return self._coerce_user_id(raw_user_id) - except Exception: - logger.warning( - "api_key_identity_cache: failed to read user", - ) - return None - - async def set_user_id( - self, - redis: RedisService, - api_key_hash: str, - user_id: str, - ttl_seconds: int, - ) -> None: - """Cache a validated API-key user ID.""" - effective_ttl_seconds: int = min(_API_KEY_MAX_TTL_SECONDS, ttl_seconds) - cache_key: str = self.get_cache_key(api_key_hash) - reverse_key: str = self.get_reverse_key(user_id) - - try: - await redis.set(cache_key, user_id, ttl=effective_ttl_seconds) - await redis.sadd(reverse_key, api_key_hash) - current_ttl_seconds = await redis.ttl(reverse_key) - if ( - current_ttl_seconds in (-2, -1) - or current_ttl_seconds < effective_ttl_seconds - ): - await redis.expire(reverse_key, effective_ttl_seconds) - except Exception: - logger.warning( - "api_key_identity_cache: failed to set user for user_id={}", - user_id, - ) - - async def invalidate_api_key( - self, - redis: RedisService, - user_id: str, - api_key_hash: str, - ) -> None: - """Delete one API-key identity cache entry.""" - try: - await redis.delete(self.get_cache_key(api_key_hash)) - await redis.srem(self.get_reverse_key(user_id), api_key_hash) - except Exception: - logger.warning( - "api_key_identity_cache: failed to invalidate identity for user_id={}", - user_id, - ) - - async def invalidate_user( - self, - redis: RedisService, - user_id: str, - ) -> None: - """Delete all API-key identity cache entries for a user.""" - try: - reverse_key: str = self.get_reverse_key(user_id) - api_key_hashes: set[object] = await redis.smembers(reverse_key) - for api_key_hash in api_key_hashes: - await redis.delete(self.get_cache_key(str(api_key_hash))) - await redis.delete(reverse_key) - except Exception: - logger.warning( - "api_key_identity_cache: failed to invalidate user_id={}", - user_id, - ) - - def _coerce_user_id(self, raw_user_id: object) -> str | None: - """Return a typed user ID from current or legacy Redis values.""" - if isinstance(raw_user_id, str): - try: - parsed_user_id = json.loads(raw_user_id) - except json.JSONDecodeError: - return raw_user_id - else: - parsed_user_id = raw_user_id - - if isinstance(parsed_user_id, str): - return parsed_user_id - - if isinstance(parsed_user_id, dict): - legacy_user_id = parsed_user_id.get("user_id") - if isinstance(legacy_user_id, str): - return legacy_user_id - - return None - - -api_key_identity_cache = APIKeyIdentityCache() diff --git a/apps/api/app/services/auth/api_key_service.py b/apps/api/app/services/auth/api_key_service.py index fc25bc12e..920751a00 100644 --- a/apps/api/app/services/auth/api_key_service.py +++ b/apps/api/app/services/auth/api_key_service.py @@ -1,12 +1,13 @@ """API key management service.""" import asyncio +import hashlib +import uuid from dataclasses import dataclass -from datetime import datetime, timezone +from datetime import datetime from typing import List, Optional from app.repositories.api_key_repository import APIKeyRepository -from app.services.auth.api_key_identity_cache import api_key_identity_cache from app.services.rate_limit.identity_cache import identity_cache from loguru import logger from sqlalchemy import select @@ -22,11 +23,6 @@ ) from shared.models.database.api_key import APIKey from shared.models.database.user_balance import UserBalance -from shared.utils.api_keys import ( - generate_api_key, - hash_api_key, - mask_api_key, -) _DEFAULT_USER_TIER: str = "free" @@ -37,7 +33,6 @@ class APIKeyIdentity: user_id: str user_tier: str - expires_at: datetime | None class APIKeyService: @@ -46,6 +41,12 @@ class APIKeyService: def __init__(self): self.repository = APIKeyRepository() + def _mask_api_key(self, api_key: str) -> str: + """Mask an API key, exposing only the first 8 and last 4 characters.""" + if not api_key or len(api_key) < 12: + return api_key + return api_key[:8] + "•" * (len(api_key) - 12) + api_key[-4:] + async def create_api_key( self, session: AsyncSession, @@ -82,10 +83,10 @@ async def create_api_key( ], ) - # 3. Generate and store a secure API key. - api_key = generate_api_key() - key_hash = hash_api_key(api_key) - key_mask = mask_api_key(api_key) + # 3. Generate a secure API key (sk_ + a 32-char UUID without hyphens). + api_key = f"sk_{str(uuid.uuid4()).replace('-', '')}" + key_hash = hashlib.sha256(api_key.encode()).hexdigest() + key_mask = self._mask_api_key(api_key) # 4. Store it in the database. api_key_record = APIKey( @@ -106,52 +107,16 @@ async def validate_api_key( self, session: AsyncSession, api_key: str ) -> Optional[str]: """Validate API key against DB, return user_id or None.""" - identity = await self.get_identity(session, api_key) + identity = await self.validate_api_key_identity(session, api_key) return identity.user_id if identity is not None else None - async def get_identity( + async def validate_api_key_identity( self, session: AsyncSession, api_key: str, ) -> Optional[APIKeyIdentity]: - """Return API-key identity, using Redis cache before DB fallback.""" - key_hash = hash_api_key(api_key) - cached_identity = await self._get_cached_identity(key_hash) - if cached_identity is not None: - return cached_identity - - identity = await self._get_database_identity(session, api_key, key_hash) - if identity is None: - return None - - await self._cache_api_key_identity( - key_hash=key_hash, - identity=identity, - ) - return identity - - async def _get_cached_identity(self, key_hash: str) -> APIKeyIdentity | None: - """Return cached API-key user identity, or None on miss/cache failure.""" - user_id = await api_key_identity_cache.get_user_id( - redis_pool_manager.get_redis_service(), - key_hash, - ) - if user_id is None: - return None - - return APIKeyIdentity( - user_id=user_id, - user_tier=await self._get_user_tier(user_id), - expires_at=None, - ) - - async def _get_database_identity( - self, - session: AsyncSession, - api_key: str, - key_hash: str, - ) -> APIKeyIdentity | None: - """Validate API key against the database.""" + """Validate API key and return the authenticated identity.""" + key_hash = hashlib.sha256(api_key.encode()).hexdigest() api_key_record = await self.repository.get_by_key_hash(session, key_hash) if not api_key_record or not api_key_record.is_valid(): @@ -159,49 +124,14 @@ async def _get_database_identity( self._schedule_last_used_update(str(api_key_record.id)) user_id = str(api_key_record.user_id) - user_tier = await self._resolve_user_tier_from_db(session, user_id) + user_tier = await self._resolve_user_tier(session, user_id) return APIKeyIdentity( user_id=user_id, user_tier=user_tier, - expires_at=api_key_record.expires_at, ) - async def _cache_api_key_identity( - self, - *, - key_hash: str, - identity: APIKeyIdentity, - ) -> None: - """Cache the validated API-key user ID for the auth layer.""" - try: - await api_key_identity_cache.set_user_id( - redis_pool_manager.get_redis_service(), - key_hash, - identity.user_id, - ttl_seconds=self._resolve_api_key_cache_ttl_seconds( - identity.expires_at - ), - ) - except Exception: - logger.warning( - "api_key_service: failed to cache identity for user_id={}", - identity.user_id, - ) - - async def _get_user_tier(self, user_id: str) -> str: - """Return user tier from rate-limit identity cache or DB fallback.""" - redis_service = redis_pool_manager.get_redis_service() - cached_identity = await identity_cache.get_user_tier(redis_service, user_id) - if cached_identity is not None: - return cached_identity["user_tier"] - - async with get_db_context() as session: - user_tier = await self._resolve_user_tier_from_db(session, user_id) - await identity_cache.set_user_tier(redis_service, user_id, user_tier) - return user_tier - - async def _resolve_user_tier_from_db( + async def _resolve_user_tier( self, session: AsyncSession, user_id: str, @@ -213,20 +143,6 @@ async def _resolve_user_tier_from_db( user_tier = result.scalar_one_or_none() return str(user_tier) if user_tier is not None else _DEFAULT_USER_TIER - def _resolve_api_key_cache_ttl_seconds(self, expires_at: datetime | None) -> int: - """Resolve cache TTL for API-key identity without exceeding key expiry.""" - max_ttl_seconds = 3600 - if expires_at is None: - return max_ttl_seconds - - expires_at_utc = expires_at - if expires_at_utc.tzinfo is None: - expires_at_utc = expires_at_utc.replace(tzinfo=timezone.utc) - - now = datetime.now(timezone.utc) - remaining_seconds = int((expires_at_utc - now).total_seconds()) - return max(1, min(max_ttl_seconds, remaining_seconds)) - async def revoke_api_key( self, session: AsyncSession, api_key_id: str, user_id: str ) -> bool: @@ -276,7 +192,7 @@ async def _invalidate_revoked_api_key_cache_best_effort( ) -> None: """Best-effort cache invalidation after a revoke has already been committed.""" try: - await api_key_identity_cache.invalidate_api_key( + await identity_cache.invalidate_apikey( redis_pool_manager.get_redis_service(), user_id, key_hash, @@ -318,10 +234,10 @@ async def regenerate_api_key( internal_message="API Key not found or does not belong to user", ) - # 2. Generate a new API key. - new_api_key = generate_api_key() - new_key_hash = hash_api_key(new_api_key) - new_key_mask = mask_api_key(new_api_key) + # 2. Generate a new API key (sk_ + a 32-char UUID without hyphens). + new_api_key = f"sk_{str(uuid.uuid4()).replace('-', '')}" + new_key_hash = hashlib.sha256(new_api_key.encode()).hexdigest() + new_key_mask = self._mask_api_key(new_api_key) # 3. Update the database record. from sqlalchemy import update @@ -339,7 +255,7 @@ async def regenerate_api_key( await session.commit() # 4. Refresh the cache. - await api_key_identity_cache.invalidate_api_key( + await identity_cache.invalidate_apikey( redis_pool_manager.get_redis_service(), user_id, api_key.key_hash, @@ -351,7 +267,7 @@ async def check_module_permission( self, session: AsyncSession, api_key: str, module: str ) -> bool: """Check whether an API key can access the requested module.""" - key_hash = hash_api_key(api_key) + key_hash = hashlib.sha256(api_key.encode()).hexdigest() api_key_record = await self.repository.get_by_key_hash(session, key_hash) if not api_key_record or not api_key_record.is_valid(): @@ -413,7 +329,7 @@ async def toggle_api_key( await session.refresh(api_key) if not api_key.is_active: - await api_key_identity_cache.invalidate_api_key( + await identity_cache.invalidate_apikey( redis_pool_manager.get_redis_service(), user_id, api_key.key_hash, diff --git a/apps/api/app/services/billing/stripe_service.py b/apps/api/app/services/billing/stripe_service.py index be1792a2e..3d10b5688 100644 --- a/apps/api/app/services/billing/stripe_service.py +++ b/apps/api/app/services/billing/stripe_service.py @@ -6,11 +6,12 @@ import stripe from app.repositories.payment_record_repository import PaymentRecordRepository from app.services.billing.price_config_service import PriceConfigService +from app.services.rate_limit.identity_cache import identity_cache from app.services.rate_limit.tier_service import TierService from sqlalchemy import func, select from sqlalchemy.ext.asyncio import AsyncSession -from shared.core.config import settings +from shared.core.config import redis_pool_manager, settings from shared.core.exceptions.domain_exceptions import ( AuthException, KnowhereException, @@ -384,6 +385,11 @@ async def _handle_checkout_completed( await db.commit() await db.refresh(payment_record) + await identity_cache.invalidate_user( + redis_pool_manager.get_redis_service(), + user_id, + ) + logger.info( f"Credits pack purchase succeeded: user_id={user_id}, credits={credits_amount}, price_id={price_id}" ) @@ -521,6 +527,11 @@ async def _handle_payment_intent_succeeded( await db.commit() await db.refresh(payment_record) + await identity_cache.invalidate_user( + redis_pool_manager.get_redis_service(), + user_id, + ) + logger.info( f"buy credits success: user_id={user_id}, credits={credits_amount}, payment_intent_id={payment_intent_id}" ) diff --git a/apps/api/app/services/guest/guest_registration_service.py b/apps/api/app/services/guest/guest_registration_service.py index 6cd16ee63..351663643 100644 --- a/apps/api/app/services/guest/guest_registration_service.py +++ b/apps/api/app/services/guest/guest_registration_service.py @@ -25,7 +25,6 @@ GuestRegisterResponse, ) from shared.services.billing.credits_service import CreditsService -from shared.utils.api_keys import generate_api_key, hash_api_key, mask_api_key _GUEST_TIER: str = "guest" _GUEST_KEY_NAME_PREFIX: str = "guest-device" @@ -154,11 +153,14 @@ async def _create_api_key_without_commit( This avoids the internal commit inside APIKeyService.create_api_key() which would make the key durable before the device row is inserted. """ + import hashlib + import uuid + from shared.models.database.api_key import APIKey - api_key = generate_api_key() - key_hash = hash_api_key(api_key) - key_mask = mask_api_key(api_key) + api_key = f"sk_{str(uuid.uuid4()).replace('-', '')}" + key_hash = hashlib.sha256(api_key.encode()).hexdigest() + key_mask = self._api_key_service._mask_api_key(api_key) api_key_record = APIKey( user_id=user_id, @@ -262,11 +264,13 @@ def _raise_existing_device_conflict(cls, device_id: str) -> NoReturn: @staticmethod async def _resolve_api_key_id(session: AsyncSession, api_key: str) -> str | None: """Resolve the DB id for a just-created API key by its hash.""" + import hashlib + from sqlalchemy import select from shared.models.database.api_key import APIKey - key_hash = hash_api_key(api_key) + key_hash = hashlib.sha256(api_key.encode()).hexdigest() result = await session.execute( select(APIKey.id).where(APIKey.key_hash == key_hash).limit(1) ) diff --git a/apps/api/app/services/rate_limit/dependencies.py b/apps/api/app/services/rate_limit/dependencies.py index 6824d66dc..c34b1c10b 100644 --- a/apps/api/app/services/rate_limit/dependencies.py +++ b/apps/api/app/services/rate_limit/dependencies.py @@ -5,9 +5,8 @@ require_billing_limits -> with_current_user -> get_current_user_id -> get_db -``with_current_user`` resolves rate-limit identity (user_id + user_tier), -caches the tier by user_id in Redis, and enforces the matched system limit -(Layer 0). +``with_current_user`` resolves identity (user_id + user_tier), caches it in +Redis, and enforces the matched system limit (Layer 0). ``require_billing_limits`` enforces billing RPM (Layer 1) when billing is enabled and yields control to the route handler. Concurrency (Layer 2) and @@ -15,7 +14,9 @@ only when billing is enabled. """ +import hashlib import math +from datetime import datetime, timezone from typing import AsyncGenerator from app.core.dependencies import get_current_user_id @@ -40,6 +41,7 @@ ) from shared.core.logging import log_context from shared.core.state_machine.states import JobStatus +from shared.models.database.api_key import APIKey from shared.models.database.job import Job from shared.models.database.user_balance import UserBalance @@ -81,6 +83,16 @@ async def _resolve_user_tier_from_db(user_id: str) -> str: return _DEFAULT_TIER +def _extract_bearer_token(authorization: str | None) -> str | None: + """Extract bearer token from Authorization header.""" + if not authorization: + return None + scheme, _, token = authorization.partition(" ") + if scheme.lower() != "bearer" or not token: + return None + return token + + def _get_route_path(request: Request) -> str: """Return the request path without the application's root_path prefix.""" scope_path: str = request.scope.get("path", request.url.path) @@ -104,6 +116,30 @@ def _get_route_limit_identifier(request: Request) -> str: return _get_route_path(request) +async def _resolve_apikey_cache_ttl_seconds(api_key_hash: str) -> int: + """Resolve cache TTL for API key identity (max 1 hour).""" + max_ttl_seconds = 3600 + try: + async with get_db_context() as db: + result = await db.execute( + select(APIKey.expires_at) + .where(APIKey.key_hash == api_key_hash) + .limit(1) + ) + expires_at = result.scalar_one_or_none() + if expires_at is None: + return max_ttl_seconds + + # APIKey.expires_at is stored as UTC-naive datetime. + if expires_at.tzinfo is None: + expires_at = expires_at.replace(tzinfo=timezone.utc) + now = datetime.now(timezone.utc) + remaining = int((expires_at - now).total_seconds()) + return max(1, min(max_ttl_seconds, remaining)) + except Exception: + return max_ttl_seconds + + # --------------------------------------------------------------------------- # with_current_user -- Layer 0 (matched system limit) # --------------------------------------------------------------------------- @@ -129,15 +165,65 @@ async def with_current_user( # -- Resolve user_tier (cache -> DB fallback) -- user_tier: str | None = getattr(request.state, "cached_user_tier", None) stashed_user_id: str | None = getattr(request.state, "user_id", None) + cached_identity_hit: bool | None = getattr( + request.state, "cached_identity_hit", None + ) - if not isinstance(user_tier, str) or stashed_user_id != user_id: + if isinstance(user_tier, str) and stashed_user_id == user_id: + if cached_identity_hit is False: + token = _extract_bearer_token(request.headers.get("authorization")) + api_key_hash = None + is_api_key_auth = isinstance(token, str) and token.startswith("sk_") + if token is not None and is_api_key_auth: + api_key_hash = hashlib.sha256(token.encode()).hexdigest() + if is_api_key_auth and api_key_hash: + try: + ttl_seconds = await _resolve_apikey_cache_ttl_seconds(api_key_hash) + await identity_cache.set_apikey_identity( + redis_service, + api_key_hash, + user_id, + user_tier, + ttl_seconds=ttl_seconds, + ) + except Exception: + logger.warning( + "rate_limit: failed to backfill API key identity cache " + "for user_id={}", + user_id, + ) + else: + token = _extract_bearer_token(request.headers.get("authorization")) + api_key_hash = None + is_api_key_auth = isinstance(token, str) and token.startswith("sk_") + if token is not None and is_api_key_auth: + api_key_hash = hashlib.sha256(token.encode()).hexdigest() + cache_key: str = ( + identity_cache._apikey_key(api_key_hash) + if is_api_key_auth and api_key_hash + else identity_cache._jwt_key(user_id) + ) try: - cached = await identity_cache.get_user_tier(redis_service, user_id) + cached: dict | None = await identity_cache.get_cached_identity( + redis_service, cache_key + ) if cached is not None: user_tier = cached.get("user_tier", _DEFAULT_TIER) else: user_tier = await _resolve_user_tier_from_db(user_id) - await identity_cache.set_user_tier(redis_service, user_id, user_tier) + if is_api_key_auth and api_key_hash: + ttl_seconds = await _resolve_apikey_cache_ttl_seconds(api_key_hash) + await identity_cache.set_apikey_identity( + redis_service, + api_key_hash, + user_id, + user_tier, + ttl_seconds=ttl_seconds, + ) + else: + await identity_cache.set_jwt_identity( + redis_service, user_id, user_tier + ) except Exception: logger.warning( "rate_limit: Redis error during identity resolution, " diff --git a/apps/api/app/services/rate_limit/identity_cache.py b/apps/api/app/services/rate_limit/identity_cache.py index 5c67b0acc..9e6ffae01 100644 --- a/apps/api/app/services/rate_limit/identity_cache.py +++ b/apps/api/app/services/rate_limit/identity_cache.py @@ -1,31 +1,69 @@ -"""Redis-backed rate-limit identity cache for user_id + user_tier.""" +""" +Redis-backed identity cache for user_id + user_tier resolution. + +Caches the mapping from authentication credentials (JWT user_id or API key hash) +to the resolved identity (user_id, user_tier) so that +tier lookups do not hit the database on every request. + +Key patterns (all prefixed with REDIS_KEY_PREFIX from config): + JWT: {REDIS_KEY_PREFIX}identity:jwt:{user_id} + API key: {REDIS_KEY_PREFIX}identity:apikey:{api_key_hash} + Reverse: {REDIS_KEY_PREFIX}identity:apikeys:{user_id} +""" import json +from typing import Optional +from app.services.rate_limit.config import REDIS_KEY_PREFIX from loguru import logger from shared.services.redis.redis_service import RedisService -_USER_TIER_TTL_SECONDS: int = 3600 +# Default TTL for JWT identity cache entries (1 hour). +_JWT_TTL_SECONDS: int = 3600 + +# Upper bound TTL for API-key identity cache entries (1 hour). +_APIKEY_MAX_TTL_SECONDS: int = 3600 class IdentityCache: - """Cache resolved rate-limit user tier by user_id.""" + """Redis-backed identity cache for user_id + user_tier resolution.""" + + # ------------------------------------------------------------------ + # Key builders + # ------------------------------------------------------------------ + + @staticmethod + def _jwt_key(user_id: str) -> str: + return f"{REDIS_KEY_PREFIX}identity:jwt:{user_id}" + + @staticmethod + def _apikey_key(api_key_hash: str) -> str: + return f"{REDIS_KEY_PREFIX}identity:apikey:{api_key_hash}" @staticmethod - def get_user_tier_key(user_id: str) -> str: - return f"identity:user:{user_id}" + def _reverse_key(user_id: str) -> str: + return f"{REDIS_KEY_PREFIX}identity:apikeys:{user_id}" - async def get_user_tier( + # ------------------------------------------------------------------ + # Read + # ------------------------------------------------------------------ + + async def get_cached_identity( self, redis: RedisService, - user_id: str, - ) -> dict[str, str] | None: + cache_key: str, + ) -> Optional[dict]: """Return cached ``{user_id, user_tier}`` or ``None`` on miss.""" - cache_key = self.get_user_tier_key(user_id) try: - raw_identity: object = await redis.get(cache_key) - return self._coerce_identity(raw_identity) + raw: Optional[str] = await redis.get(cache_key) + if raw is None: + return None + # RedisService.get already attempts JSON parse, but the + # value may come back as a dict directly. + if isinstance(raw, dict): + return raw + return json.loads(raw) except Exception: logger.warning( "identity_cache: failed to read cache_key={}", @@ -33,52 +71,121 @@ async def get_user_tier( ) return None - async def set_user_tier( + # ------------------------------------------------------------------ + # Write -- JWT + # ------------------------------------------------------------------ + + async def set_jwt_identity( self, redis: RedisService, user_id: str, user_tier: str, ) -> None: - """Cache rate-limit identity for a user.""" - key: str = self.get_user_tier_key(user_id) - payload: dict[str, str] = {"user_id": user_id, "user_tier": user_tier} + """Cache identity for a JWT-authenticated user (1 hr TTL).""" + key: str = self._jwt_key(user_id) + payload: dict = {"user_id": user_id, "user_tier": user_tier} try: - await redis.set(key, payload, ttl=_USER_TIER_TTL_SECONDS) + await redis.set(key, payload, ttl=_JWT_TTL_SECONDS) except Exception: logger.warning( - "identity_cache: failed to set user tier for user_id={}", + "identity_cache: failed to set jwt identity user_id={}", user_id, ) - async def invalidate_user( + # ------------------------------------------------------------------ + # Write -- API key + # ------------------------------------------------------------------ + + async def set_apikey_identity( self, redis: RedisService, + api_key_hash: str, user_id: str, + user_tier: str, + ttl_seconds: int, ) -> None: - """Delete cached rate-limit identity for a user.""" + """Cache identity for an API-key-authenticated user. + + TTL is ``min(APIKEY_MAX_TTL, api_key_remaining_ttl)`` so the + cache never outlives the key itself. Also maintains a reverse + index (SET) of all cached API-key hashes per user for bulk + invalidation. + """ + effective_ttl: int = min(_APIKEY_MAX_TTL_SECONDS, ttl_seconds) + key: str = self._apikey_key(api_key_hash) + payload: dict = {"user_id": user_id, "user_tier": user_tier} try: - await redis.delete(self.get_user_tier_key(user_id)) + await redis.set(key, payload, ttl=effective_ttl) + # Maintain reverse index so invalidate_user can find all + # API-key cache entries belonging to this user. + reverse_key: str = self._reverse_key(user_id) + await redis.sadd(reverse_key, api_key_hash) + # Keep reverse index TTL at least as long as the longest + # surviving API-key cache entry for this user. + current_ttl = await redis.ttl(reverse_key) + if current_ttl in (-2, -1) or current_ttl < effective_ttl: + await redis.expire(reverse_key, effective_ttl) except Exception: logger.warning( - "identity_cache: failed to invalidate user_id={}", + "identity_cache: failed to set apikey identity " + "api_key_hash={}, user_id={}", + api_key_hash, user_id, ) - def _coerce_identity(self, raw_identity: object) -> dict[str, str] | None: - """Return a typed identity payload from a Redis value.""" - parsed_identity = raw_identity - if isinstance(raw_identity, str): - parsed_identity = json.loads(raw_identity) + # ------------------------------------------------------------------ + # Invalidation + # ------------------------------------------------------------------ - if not isinstance(parsed_identity, dict): - return None + async def invalidate_user( + self, + redis: RedisService, + user_id: str, + ) -> None: + """Full invalidation: JWT cache + all API-key caches + reverse index.""" + try: + # 1. Delete JWT cache + jwt_key: str = self._jwt_key(user_id) + await redis.delete(jwt_key) - user_id = parsed_identity.get("user_id") - user_tier = parsed_identity.get("user_tier") - if not isinstance(user_id, str) or not isinstance(user_tier, str): - return None + # 2. Collect all cached API-key hashes from reverse index + reverse_key: str = self._reverse_key(user_id) + api_key_hashes: set = await redis.smembers(reverse_key) + + # 3. Delete each API-key cache entry + for api_key_hash in api_key_hashes: + apikey_key: str = self._apikey_key(str(api_key_hash)) + await redis.delete(apikey_key) + + # 4. Delete the reverse index itself + await redis.delete(reverse_key) + except Exception: + logger.warning( + "identity_cache: failed to invalidate user_id={}", + user_id, + ) - return {"user_id": user_id, "user_tier": user_tier} + async def invalidate_apikey( + self, + redis: RedisService, + user_id: str, + api_key_hash: str, + ) -> None: + """Delete a single API-key cache entry and remove from reverse index.""" + try: + apikey_key: str = self._apikey_key(api_key_hash) + await redis.delete(apikey_key) + + reverse_key: str = self._reverse_key(user_id) + await redis.srem(reverse_key, api_key_hash) + except Exception: + logger.warning( + "identity_cache: failed to invalidate apikey " + "api_key_hash={}, user_id={}", + api_key_hash, + user_id, + ) +# Module-level singleton so callers can import directly. identity_cache = IdentityCache() diff --git a/apps/api/app/services/rate_limit/tier_service.py b/apps/api/app/services/rate_limit/tier_service.py index e86a4f782..a9d73a13e 100644 --- a/apps/api/app/services/rate_limit/tier_service.py +++ b/apps/api/app/services/rate_limit/tier_service.py @@ -7,12 +7,10 @@ from app.services.rate_limit.config import RateLimitConfig from app.services.rate_limit.data_structures import TierLimits -from app.services.rate_limit.identity_cache import identity_cache from loguru import logger from sqlalchemy import func, select, update from sqlalchemy.ext.asyncio import AsyncSession -from shared.core.config import redis_pool_manager from shared.models.database.payment_record import PaymentRecord from shared.models.database.tier_limit import TierLimit from shared.models.database.user_balance import UserBalance @@ -63,7 +61,6 @@ async def refresh_tier(user_id: str, session: AsyncSession) -> str: .values(user_tier=new_tier) ) await session.execute(stmt_update) - await TierService._invalidate_identity_caches(user_id) logger.info( "Tier refreshed: user_id=%s total_micro=%d new_tier=%s", @@ -81,15 +78,3 @@ def get_limits(user_tier: str) -> Optional[TierLimits]: """ config = RateLimitConfig.get_instance() return config.tier_map.get(user_tier) - - @staticmethod - async def _invalidate_identity_caches(user_id: str) -> None: - """Invalidate identity caches that include user_tier.""" - try: - redis_service = redis_pool_manager.get_redis_service() - await identity_cache.invalidate_user(redis_service, user_id) - except Exception: - logger.warning( - "Tier refresh cache invalidation failed for user_id={}", - user_id, - ) diff --git a/apps/api/scripts/bootstrap_local_dev.py b/apps/api/scripts/bootstrap_local_dev.py index 7fbea96b6..779dee503 100644 --- a/apps/api/scripts/bootstrap_local_dev.py +++ b/apps/api/scripts/bootstrap_local_dev.py @@ -45,7 +45,7 @@ async def _run(mode: str) -> int: def _print_profile() -> None: - profile = LocalDevelopmentBootstrapService.get_local_developer_auth_profile() + profile = LocalDevelopmentBootstrapService.get_local_developer_profile() print(f"user_id={profile['user_id']}") print(f"name={profile['name']}") print(f"email={profile['email']}") diff --git a/apps/api/scripts/init_user.py b/apps/api/scripts/init_user.py index c7442fc85..2c156ed0e 100644 --- a/apps/api/scripts/init_user.py +++ b/apps/api/scripts/init_user.py @@ -2,9 +2,10 @@ import argparse import asyncio +import hashlib import os +import secrets import sys -from pathlib import Path from uuid import uuid4 from sqlalchemy import select @@ -19,7 +20,6 @@ from shared.models.database.user import User from shared.models.database.user_balance import UserBalance from shared.services.auth.user_table_bootstrap import ensure_better_auth_user_table -from shared.utils.api_keys import generate_api_key, hash_api_key, mask_api_key _DEFAULT_API_KEY_NAME: str = "standalone-api-key" _DEFAULT_USER_TIER: str = "free" @@ -43,11 +43,6 @@ def _build_parser() -> argparse.ArgumentParser: default=_DEFAULT_USER_TIER, help="Compatibility user tier to store in user_balances.", ) - parser.add_argument( - "--api-key-output-file", - default="", - help="Optional file path for the generated API key. The file is created with 0600 permissions.", - ) return parser @@ -129,17 +124,14 @@ async def _resolve_key_name( return f"{key_name}-{suffix}" -def _write_api_key_file(path_value: str, api_key: str) -> Path: - output_path = Path(path_value).expanduser() - output_path.parent.mkdir(parents=True, exist_ok=True) - file_descriptor = os.open( - output_path, - os.O_WRONLY | os.O_CREAT | os.O_TRUNC, - 0o600, - ) - with os.fdopen(file_descriptor, "w", encoding="utf-8") as output_file: - output_file.write(f"{api_key}\n") - return output_path +def _generate_api_key() -> str: + return f"sk_kn_{secrets.token_hex(16)}" + + +def _mask_api_key(api_key: str) -> str: + if len(api_key) < 12: + return api_key + return api_key[:8] + "•" * (len(api_key) - 12) + api_key[-4:] async def _create_api_key( @@ -148,12 +140,12 @@ async def _create_api_key( user_id: str, key_name: str, ) -> str: - api_key = generate_api_key() + api_key = _generate_api_key() session.add( APIKey( user_id=user_id, - key_hash=hash_api_key(api_key), - key_mask=mask_api_key(api_key), + key_hash=hashlib.sha256(api_key.encode()).hexdigest(), + key_mask=_mask_api_key(api_key), name=key_name, enabled_modules=["all"], ) @@ -187,15 +179,8 @@ async def _run(args: argparse.Namespace) -> int: print(f"user_id={user.id}") print(f"email={user.email}") - print("credential_name_created=true") - print("api_key_created=true") - print("credential_hidden=true") - output_path_value = str(args.api_key_output_file).strip() - if output_path_value: - output_path = _write_api_key_file(output_path_value, api_key) - print(f"credential_output_file={output_path}") - else: - print("credential_output_file=") + print(f"api_key_name={key_name}") + print(f"api_key={api_key}") return 0 diff --git a/apps/api/scripts/local_dev_bootstrap_service.py b/apps/api/scripts/local_dev_bootstrap_service.py index 7a98088ab..758ecaab6 100644 --- a/apps/api/scripts/local_dev_bootstrap_service.py +++ b/apps/api/scripts/local_dev_bootstrap_service.py @@ -1,5 +1,6 @@ from __future__ import annotations +import hashlib from datetime import datetime, timezone from sqlalchemy.ext.asyncio import AsyncSession @@ -12,7 +13,6 @@ from shared.models.database.user import User from shared.models.database.user_balance import UserBalance from shared.services.auth.user_table_bootstrap import ensure_better_auth_user_table -from shared.utils.api_keys import hash_api_key, mask_api_key class LocalDevelopmentBootstrapService: @@ -52,23 +52,16 @@ async def seed_local_developer(self, session: AsyncSession) -> None: @classmethod def get_local_developer_profile(cls) -> dict[str, str | int]: - """Expose deterministic local developer profile details for local tooling.""" - profile: dict[str, str | int] = { + """Expose deterministic local developer credentials for local tooling.""" + return { "user_id": cls.LOCAL_DEV_USER_ID, "name": cls.LOCAL_DEV_USER_NAME, "email": cls.LOCAL_DEV_USER_EMAIL, "tier": cls.LOCAL_DEV_TIER, + "api_key": cls.LOCAL_DEV_API_KEY, "credits_balance": cls.LOCAL_DEV_CREDITS_BALANCE, "lifetime_billing_micro": cls.LOCAL_DEV_LIFETIME_BILLING_MICRO, } - return profile - - @classmethod - def get_local_developer_auth_profile(cls) -> dict[str, str | int]: - """Expose deterministic local developer auth details for contract tests.""" - auth_profile = cls.get_local_developer_profile() - auth_profile["api_key"] = cls.LOCAL_DEV_API_KEY - return auth_profile async def _upsert_user(self, session: AsyncSession) -> None: user = await session.get(User, self.LOCAL_DEV_USER_ID) @@ -155,8 +148,8 @@ async def _upsert_credits_transaction(self, session: AsyncSession) -> None: async def _upsert_api_key(self, session: AsyncSession) -> None: api_key = await session.get(APIKey, self.LOCAL_DEV_API_KEY_ID) - key_hash = hash_api_key(self.LOCAL_DEV_API_KEY) - key_mask = mask_api_key(self.LOCAL_DEV_API_KEY) + key_hash = hashlib.sha256(self.LOCAL_DEV_API_KEY.encode()).hexdigest() + key_mask = self._mask_api_key(self.LOCAL_DEV_API_KEY) if api_key is None: session.add( @@ -178,6 +171,12 @@ async def _upsert_api_key(self, session: AsyncSession) -> None: api_key.enabled_modules = ["all"] api_key.is_active = True + @staticmethod + def _mask_api_key(api_key: str) -> str: + if len(api_key) < 12: + return api_key + return api_key[:8] + "•" * (len(api_key) - 12) + api_key[-4:] + @staticmethod def _utc_now() -> datetime: return datetime.now(timezone.utc).replace(tzinfo=None) diff --git a/apps/api/tests/contract/test_api_key_contract.py b/apps/api/tests/contract/test_api_key_contract.py index ec55ff866..21b693fb3 100644 --- a/apps/api/tests/contract/test_api_key_contract.py +++ b/apps/api/tests/contract/test_api_key_contract.py @@ -6,9 +6,6 @@ import pytest from httpx import AsyncClient -from shared.utils.api_keys import hash_api_key -from tests.support.contract_database import ContractDatabase - @pytest.mark.asyncio async def test_should_revoke_a_created_api_key_through_http_only( @@ -81,37 +78,6 @@ async def test_should_revoke_a_created_api_key_through_http_only( assert "details" not in error -@pytest.mark.asyncio -async def test_should_accept_an_active_sha256_api_key_hash( - api_client_factory: Callable[[], AbstractAsyncContextManager[AsyncClient]], -) -> None: - user_id = f"sha256-user-{uuid4().hex[:12]}" - raw_api_key = f"sk_sha256_{uuid4().hex}" - key_hash = hash_api_key(raw_api_key) - - async with api_client_factory() as api_client: - await ContractDatabase.insert_authenticated_user( - user_id=user_id, - api_key=raw_api_key, - user_tier="tier_5", - ) - await ContractDatabase.execute( - """ - UPDATE api_keys - SET key_hash = :key_hash - WHERE user_id = :user_id - """, - { - "key_hash": key_hash, - "user_id": user_id, - }, - ) - api_client.headers.update({"Authorization": f"Bearer {raw_api_key}"}) - response = await api_client.get("/api/v1/jobs") - - assert response.status_code == 200 - - @pytest.mark.asyncio async def test_should_regenerate_an_api_key_and_invalidate_the_previous_raw_key( developer_api_client_factory: Callable[ diff --git a/apps/api/tests/support/contract_database.py b/apps/api/tests/support/contract_database.py index 1a6d34bbf..681417188 100644 --- a/apps/api/tests/support/contract_database.py +++ b/apps/api/tests/support/contract_database.py @@ -1,5 +1,6 @@ from __future__ import annotations +import hashlib import json from datetime import datetime, timezone from typing import Any @@ -9,7 +10,6 @@ from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine from shared.testing.contract_runtime import get_contract_database_url -from shared.utils.api_keys import hash_api_key, mask_api_key async def _create_contract_engine() -> AsyncEngine: @@ -139,7 +139,7 @@ async def insert_authenticated_user( ) timestamp = _utc_now() - api_key_hash = hash_api_key(api_key) + api_key_hash = hashlib.sha256(api_key.encode()).hexdigest() api_key_id = f"key_{uuid4().hex[:12]}" await cls.execute( @@ -168,7 +168,7 @@ async def insert_authenticated_user( "id": api_key_id, "user_id": user_id, "key_hash": api_key_hash, - "key_mask": mask_api_key(api_key), + "key_mask": f"{api_key[:8]}...{api_key[-4:]}", "name": f"Contract API Key {user_id}", "enabled_modules": json.dumps(enabled_modules or ["all"]), "is_active": True, diff --git a/deploy/local-dev/README.md b/deploy/local-dev/README.md index 9b102c66d..e481b8033 100644 --- a/deploy/local-dev/README.md +++ b/deploy/local-dev/README.md @@ -45,7 +45,7 @@ Deterministic local developer account: - `user_id`: `local-dev-user` - `email`: `local-dev-user@knowhere.local` - `tier`: `tier_5` -- `local_developer_key_seeded`: `true` +- `api_key`: `local_dev_demo_key_tier5_full_access` ## Verify the Local API diff --git a/packages/shared-python/shared/testing/contract_runtime.py b/packages/shared-python/shared/testing/contract_runtime.py index 5ea5cfd29..8c816ce79 100644 --- a/packages/shared-python/shared/testing/contract_runtime.py +++ b/packages/shared-python/shared/testing/contract_runtime.py @@ -657,7 +657,7 @@ async def seed_contract_developer() -> dict[str, str | int]: finally: await engine.dispose() - return bootstrap_module.LocalDevelopmentBootstrapService.get_local_developer_auth_profile() + return bootstrap_module.LocalDevelopmentBootstrapService.get_local_developer_profile() async def reset_contract_database() -> None: diff --git a/packages/shared-python/shared/tests/utils/test_api_keys.py b/packages/shared-python/shared/tests/utils/test_api_keys.py deleted file mode 100644 index c4984474d..000000000 --- a/packages/shared-python/shared/tests/utils/test_api_keys.py +++ /dev/null @@ -1,34 +0,0 @@ -from shared.utils.api_keys import ( - API_KEY_PREFIX, - generate_api_key, - hash_api_key, - is_api_key_token, - mask_api_key, -) - - -def test_generate_api_key_should_use_api_key_prefix_and_random_secret() -> None: - first_api_key = generate_api_key() - second_api_key = generate_api_key() - - assert first_api_key.startswith(API_KEY_PREFIX) - assert second_api_key.startswith(API_KEY_PREFIX) - assert first_api_key != second_api_key - assert len(first_api_key) > len(API_KEY_PREFIX) + 32 - - -def test_hash_api_key_should_return_deterministic_sha256_lookup_hash() -> None: - api_key = "sk_contract_test_secret" - - assert hash_api_key(api_key) == hash_api_key(api_key) - assert len(hash_api_key(api_key)) == 64 - - -def test_mask_api_key_should_hide_middle_characters() -> None: - assert mask_api_key("sk_1234567890abcdef") == "sk_12345•••••••cdef" - - -def test_is_api_key_token_should_match_only_api_key_prefix() -> None: - assert is_api_key_token("sk_test") is True - assert is_api_key_token("jwt_test") is False - assert is_api_key_token(None) is False diff --git a/packages/shared-python/shared/utils/api_keys.py b/packages/shared-python/shared/utils/api_keys.py deleted file mode 100644 index 71b5f8add..000000000 --- a/packages/shared-python/shared/utils/api_keys.py +++ /dev/null @@ -1,30 +0,0 @@ -"""API key generation, masking, and hashing helpers.""" - -from hashlib import sha256 -from secrets import token_urlsafe -from typing import TypeGuard - -API_KEY_PREFIX = "sk_" -API_KEY_RANDOM_BYTES = 32 - - -def hash_api_key(api_key: str) -> str: - """Return a deterministic SHA-256 digest for API key lookup.""" - return sha256(api_key.encode("utf-8")).hexdigest() - - -def generate_api_key() -> str: - """Generate a new plaintext API key with cryptographic randomness.""" - return f"{API_KEY_PREFIX}{token_urlsafe(API_KEY_RANDOM_BYTES)}" - - -def mask_api_key(api_key: str) -> str: - """Mask an API key, exposing only the first 8 and last 4 characters.""" - if len(api_key) < 12: - return api_key - return api_key[:8] + "•" * (len(api_key) - 12) + api_key[-4:] - - -def is_api_key_token(token: object) -> TypeGuard[str]: - """Return whether a bearer token has the API-key prefix.""" - return isinstance(token, str) and token.startswith(API_KEY_PREFIX) From 5cd99c6cd50ac86bf15edddc06a5aeb86d5adc8f Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 04:44:43 +0000 Subject: [PATCH 17/32] refactor: split api key and tier identity caches --- apps/api/app/core/dependencies.py | 38 +--- .../services/auth/api_key_identity_cache.py | 124 +++++++++++++ apps/api/app/services/auth/api_key_service.py | 151 +++++++++++---- .../app/services/billing/stripe_service.py | 13 +- .../guest/guest_registration_service.py | 14 +- .../app/services/rate_limit/dependencies.py | 124 ++----------- .../app/services/rate_limit/identity_cache.py | 175 ++++-------------- .../app/services/rate_limit/tier_service.py | 17 ++ .../contract/test_identity_cache_contract.py | 117 ++++++++++++ .../shared/tests/utils/test_api_keys.py | 34 ++++ .../shared-python/shared/utils/api_keys.py | 30 +++ 11 files changed, 498 insertions(+), 339 deletions(-) create mode 100644 apps/api/app/services/auth/api_key_identity_cache.py create mode 100644 apps/api/tests/contract/test_identity_cache_contract.py create mode 100644 packages/shared-python/shared/tests/utils/test_api_keys.py create mode 100644 packages/shared-python/shared/utils/api_keys.py diff --git a/apps/api/app/core/dependencies.py b/apps/api/app/core/dependencies.py index ac349e34f..02038f5c9 100644 --- a/apps/api/app/core/dependencies.py +++ b/apps/api/app/core/dependencies.py @@ -1,4 +1,3 @@ -import hashlib import threading from datetime import timedelta from fnmatch import fnmatch @@ -6,18 +5,18 @@ import jwt from app.services.auth.api_key_service import APIKeyService -from app.services.rate_limit.identity_cache import identity_cache from fastapi import Depends, Header, Request from jwt import PyJWKClient from loguru import logger from sqlalchemy.ext.asyncio import AsyncSession -from shared.core.config import redis_pool_manager, settings +from shared.core.config import settings from shared.core.database import get_db from shared.core.exceptions.domain_exceptions import ( AuthException, PermissionDeniedException, ) +from shared.utils.api_keys import is_api_key_token # Standard JWKS endpoint path (fixed, following OpenID Connect convention) JWKS_ENDPOINT_PATH = "/api/auth/jwks" @@ -184,39 +183,14 @@ async def get_current_user_id( route_path = _get_route_path(request) # Mode 1: API Key verification (for external clients) - if token.startswith("sk_"): - # Check identity cache first — skip DB on cache hit - api_key_hash = hashlib.sha256(token.encode()).hexdigest() - try: - cached = await identity_cache.get_cached_identity( - redis_pool_manager.get_redis_service(), - identity_cache._apikey_key(api_key_hash), - ) - if cached is not None: - cached_user_id = cached.get("user_id") - cached_user_tier = cached.get("user_tier") - if cached_user_id and isinstance(cached_user_tier, str): - request.state.cached_user_tier = cached_user_tier - request.state.cached_identity_hit = True - request.state.user_id = cached_user_id - _enforce_guest_api_key_scope(route_path, cached_user_tier) - return cached_user_id - except PermissionDeniedException: - raise - except Exception: - pass # Fall through to DB validation - - # Cache miss — validate via DB + if is_api_key_token(token): api_key_service = APIKeyService() - identity = await api_key_service.validate_api_key_identity(db, token) + identity = await api_key_service.get_identity(db, token) if identity: - request.state.cached_user_tier = identity.user_tier - request.state.cached_identity_hit = False - request.state.user_id = identity.user_id _enforce_guest_api_key_scope(route_path, identity.user_tier) return identity.user_id - else: - raise AuthException(user_message="Invalid API Key") + + raise AuthException(user_message="Invalid API Key") # Mode 2: JWT verification (for Dashboard/Internal) return decode_jwt_token(token) diff --git a/apps/api/app/services/auth/api_key_identity_cache.py b/apps/api/app/services/auth/api_key_identity_cache.py new file mode 100644 index 000000000..95308853c --- /dev/null +++ b/apps/api/app/services/auth/api_key_identity_cache.py @@ -0,0 +1,124 @@ +"""Redis-backed API-key authentication user cache.""" + +from __future__ import annotations + +import json +from typing import TYPE_CHECKING + +from loguru import logger + +if TYPE_CHECKING: + from shared.services.redis.redis_service import RedisService + +_API_KEY_MAX_TTL_SECONDS: int = 3600 + + +class APIKeyIdentityCache: + """Cache validated API-key user IDs by API-key lookup hash.""" + + @staticmethod + def get_cache_key(api_key_hash: str) -> str: + """Return the Redis key for an API-key hash.""" + return f"identity:apikey:{api_key_hash}" + + @staticmethod + def get_reverse_key(user_id: str) -> str: + """Return the reverse-index Redis key for a user.""" + return f"identity:apikeys:{user_id}" + + async def get_user_id( + self, + redis: RedisService, + api_key_hash: str, + ) -> str | None: + """Return cached user_id for an API key.""" + try: + raw_user_id: object = await redis.get(self.get_cache_key(api_key_hash)) + return self._coerce_user_id(raw_user_id) + except Exception: + logger.warning("api_key_identity_cache: failed to read user") + return None + + async def set_user_id( + self, + redis: RedisService, + api_key_hash: str, + user_id: str, + ttl_seconds: int, + ) -> None: + """Cache a validated API-key user ID.""" + effective_ttl_seconds: int = min(_API_KEY_MAX_TTL_SECONDS, ttl_seconds) + cache_key: str = self.get_cache_key(api_key_hash) + reverse_key: str = self.get_reverse_key(user_id) + + try: + await redis.set(cache_key, user_id, ttl=effective_ttl_seconds) + await redis.sadd(reverse_key, api_key_hash) + current_ttl_seconds: int = await redis.ttl(reverse_key) + if ( + current_ttl_seconds in (-2, -1) + or current_ttl_seconds < effective_ttl_seconds + ): + await redis.expire(reverse_key, effective_ttl_seconds) + except Exception: + logger.warning( + "api_key_identity_cache: failed to set user for user_id={}", + user_id, + ) + + async def invalidate_api_key( + self, + redis: RedisService, + user_id: str, + api_key_hash: str, + ) -> None: + """Delete one API-key identity cache entry.""" + try: + await redis.delete(self.get_cache_key(api_key_hash)) + await redis.srem(self.get_reverse_key(user_id), api_key_hash) + except Exception: + logger.warning( + "api_key_identity_cache: failed to invalidate identity for user_id={}", + user_id, + ) + + async def invalidate_user( + self, + redis: RedisService, + user_id: str, + ) -> None: + """Delete all API-key identity cache entries for a user.""" + try: + reverse_key: str = self.get_reverse_key(user_id) + api_key_hashes: set[object] = await redis.smembers(reverse_key) + for api_key_hash in api_key_hashes: + await redis.delete(self.get_cache_key(str(api_key_hash))) + await redis.delete(reverse_key) + except Exception: + logger.warning( + "api_key_identity_cache: failed to invalidate user_id={}", + user_id, + ) + + def _coerce_user_id(self, raw_user_id: object) -> str | None: + """Return a typed user ID from current or legacy Redis values.""" + if isinstance(raw_user_id, str): + try: + parsed_user_id: object = json.loads(raw_user_id) + except json.JSONDecodeError: + return raw_user_id + else: + parsed_user_id = raw_user_id + + if isinstance(parsed_user_id, str): + return parsed_user_id + + if isinstance(parsed_user_id, dict): + legacy_user_id: object = parsed_user_id.get("user_id") + if isinstance(legacy_user_id, str): + return legacy_user_id + + return None + + +api_key_identity_cache = APIKeyIdentityCache() diff --git a/apps/api/app/services/auth/api_key_service.py b/apps/api/app/services/auth/api_key_service.py index 920751a00..b79baa3ae 100644 --- a/apps/api/app/services/auth/api_key_service.py +++ b/apps/api/app/services/auth/api_key_service.py @@ -1,16 +1,15 @@ """API key management service.""" import asyncio -import hashlib -import uuid from dataclasses import dataclass -from datetime import datetime +from datetime import datetime, timezone from typing import List, Optional from app.repositories.api_key_repository import APIKeyRepository +from app.services.auth.api_key_identity_cache import api_key_identity_cache from app.services.rate_limit.identity_cache import identity_cache from loguru import logger -from sqlalchemy import select +from sqlalchemy import select, update from sqlalchemy.ext.asyncio import AsyncSession from shared.core.config import redis_pool_manager @@ -23,8 +22,10 @@ ) from shared.models.database.api_key import APIKey from shared.models.database.user_balance import UserBalance +from shared.utils.api_keys import generate_api_key, hash_api_key, mask_api_key _DEFAULT_USER_TIER: str = "free" +_API_KEY_MAX_CACHE_TTL_SECONDS: int = 3600 @dataclass(frozen=True) @@ -33,6 +34,7 @@ class APIKeyIdentity: user_id: str user_tier: str + expires_at: datetime | None class APIKeyService: @@ -43,9 +45,7 @@ def __init__(self): def _mask_api_key(self, api_key: str) -> str: """Mask an API key, exposing only the first 8 and last 4 characters.""" - if not api_key or len(api_key) < 12: - return api_key - return api_key[:8] + "•" * (len(api_key) - 12) + api_key[-4:] + return mask_api_key(api_key) async def create_api_key( self, @@ -56,9 +56,8 @@ async def create_api_key( expires_at: Optional[datetime] = None, ) -> str: """Create an API key.""" - # 1. Enforce the per-user API key limit. key_count = await self.repository.count_by_user(session, user_id) - if key_count >= 10: # Limit each user to at most 10 API keys. + if key_count >= 10: raise ValidationException( user_message="Maximum API Key limit reached (10)", violations=[ @@ -83,19 +82,16 @@ async def create_api_key( ], ) - # 3. Generate a secure API key (sk_ + a 32-char UUID without hyphens). - api_key = f"sk_{str(uuid.uuid4()).replace('-', '')}" - key_hash = hashlib.sha256(api_key.encode()).hexdigest() - key_mask = self._mask_api_key(api_key) + api_key = generate_api_key() + key_hash = hash_api_key(api_key) + key_mask = mask_api_key(api_key) - # 4. Store it in the database. api_key_record = APIKey( user_id=user_id, key_hash=key_hash, key_mask=key_mask, name=name, - enabled_modules=enabled_modules - or ["all"], # Enable all modules by default. + enabled_modules=enabled_modules or ["all"], expires_at=expires_at, ) @@ -107,7 +103,7 @@ async def validate_api_key( self, session: AsyncSession, api_key: str ) -> Optional[str]: """Validate API key against DB, return user_id or None.""" - identity = await self.validate_api_key_identity(session, api_key) + identity = await self.get_identity(session, api_key) return identity.user_id if identity is not None else None async def validate_api_key_identity( @@ -116,7 +112,47 @@ async def validate_api_key_identity( api_key: str, ) -> Optional[APIKeyIdentity]: """Validate API key and return the authenticated identity.""" - key_hash = hashlib.sha256(api_key.encode()).hexdigest() + return await self.get_identity(session, api_key) + + async def get_identity( + self, + session: AsyncSession, + api_key: str, + ) -> APIKeyIdentity | None: + """Return API-key identity, using auth cache before DB fallback.""" + key_hash: str = hash_api_key(api_key) + cached_identity = await self._get_cached_identity(key_hash) + if cached_identity is not None: + return cached_identity + + identity = await self._get_database_identity(session, key_hash) + if identity is None: + return None + + await self._cache_api_key_identity(key_hash=key_hash, identity=identity) + return identity + + async def _get_cached_identity(self, key_hash: str) -> APIKeyIdentity | None: + """Return cached API-key user identity, or None on miss/cache failure.""" + user_id = await api_key_identity_cache.get_user_id( + redis_pool_manager.get_redis_service(), + key_hash, + ) + if user_id is None: + return None + + return APIKeyIdentity( + user_id=user_id, + user_tier=await self._get_user_tier(user_id), + expires_at=None, + ) + + async def _get_database_identity( + self, + session: AsyncSession, + key_hash: str, + ) -> APIKeyIdentity | None: + """Validate API key against the database.""" api_key_record = await self.repository.get_by_key_hash(session, key_hash) if not api_key_record or not api_key_record.is_valid(): @@ -124,14 +160,49 @@ async def validate_api_key_identity( self._schedule_last_used_update(str(api_key_record.id)) user_id = str(api_key_record.user_id) - user_tier = await self._resolve_user_tier(session, user_id) + user_tier = await self._resolve_user_tier_from_db(session, user_id) return APIKeyIdentity( user_id=user_id, user_tier=user_tier, + expires_at=api_key_record.expires_at, ) - async def _resolve_user_tier( + async def _cache_api_key_identity( + self, + *, + key_hash: str, + identity: APIKeyIdentity, + ) -> None: + """Cache the validated API-key user ID for the auth layer.""" + try: + await api_key_identity_cache.set_user_id( + redis_pool_manager.get_redis_service(), + key_hash, + identity.user_id, + ttl_seconds=self._resolve_api_key_cache_ttl_seconds( + identity.expires_at + ), + ) + except Exception: + logger.warning( + "api_key_service: failed to cache identity for user_id={}", + identity.user_id, + ) + + async def _get_user_tier(self, user_id: str) -> str: + """Return user tier from rate-limit cache or DB fallback.""" + redis_service = redis_pool_manager.get_redis_service() + cached_identity = await identity_cache.get_user_tier(redis_service, user_id) + if cached_identity is not None: + return cached_identity["user_tier"] + + async with get_db_context() as session: + user_tier = await self._resolve_user_tier_from_db(session, user_id) + await identity_cache.set_user_tier(redis_service, user_id, user_tier) + return user_tier + + async def _resolve_user_tier_from_db( self, session: AsyncSession, user_id: str, @@ -143,13 +214,25 @@ async def _resolve_user_tier( user_tier = result.scalar_one_or_none() return str(user_tier) if user_tier is not None else _DEFAULT_USER_TIER + def _resolve_api_key_cache_ttl_seconds(self, expires_at: datetime | None) -> int: + """Resolve cache TTL for API-key identity without exceeding key expiry.""" + if expires_at is None: + return _API_KEY_MAX_CACHE_TTL_SECONDS + + expires_at_utc = expires_at + if expires_at_utc.tzinfo is None: + expires_at_utc = expires_at_utc.replace(tzinfo=timezone.utc) + + now = datetime.now(timezone.utc) + remaining_seconds = int((expires_at_utc - now).total_seconds()) + return max(1, min(_API_KEY_MAX_CACHE_TTL_SECONDS, remaining_seconds)) + async def revoke_api_key( self, session: AsyncSession, api_key_id: str, user_id: str ) -> bool: """Revoke an API key by deleting it directly.""" logger.info(f"Revoking API key: api_key_id={api_key_id}, user_id={user_id}") - # 1. Verify that the API key belongs to the user. api_key = await self.repository.get_by_id(session, api_key_id) if not api_key: @@ -170,11 +253,9 @@ async def revoke_api_key( internal_message="API Key not found or does not belong to user", ) - # 2. Delete the API key directly. success = await self.repository.delete_by_id(session, api_key_id) logger.info(f"Delete result: {success}") - # 3. Commit the transaction. if success: await session.commit() logger.info("Transaction committed") @@ -192,7 +273,7 @@ async def _invalidate_revoked_api_key_cache_best_effort( ) -> None: """Best-effort cache invalidation after a revoke has already been committed.""" try: - await identity_cache.invalidate_apikey( + await api_key_identity_cache.invalidate_api_key( redis_pool_manager.get_redis_service(), user_id, key_hash, @@ -212,7 +293,7 @@ async def list_user_api_keys( "id": str(api_key.id), "name": api_key.name, "api_key": api_key.key_mask - or f"sk_{api_key.id[:8]}••••••••••••••••••••••••••••••••••••••••", # Return the masked API key. + or f"sk_{api_key.id[:8]}••••••••••••••••••••••••••••••••••••••••", "enabled_modules": api_key.enabled_modules, "is_active": api_key.is_active, "created_at": api_key.created_at, @@ -234,15 +315,9 @@ async def regenerate_api_key( internal_message="API Key not found or does not belong to user", ) - # 2. Generate a new API key (sk_ + a 32-char UUID without hyphens). - new_api_key = f"sk_{str(uuid.uuid4()).replace('-', '')}" - new_key_hash = hashlib.sha256(new_api_key.encode()).hexdigest() - new_key_mask = self._mask_api_key(new_api_key) - - # 3. Update the database record. - from sqlalchemy import update - - from shared.models.database.api_key import APIKey + new_api_key = generate_api_key() + new_key_hash = hash_api_key(new_api_key) + new_key_mask = mask_api_key(new_api_key) await session.execute( update(APIKey) @@ -254,8 +329,7 @@ async def regenerate_api_key( ) await session.commit() - # 4. Refresh the cache. - await identity_cache.invalidate_apikey( + await api_key_identity_cache.invalidate_api_key( redis_pool_manager.get_redis_service(), user_id, api_key.key_hash, @@ -267,13 +341,12 @@ async def check_module_permission( self, session: AsyncSession, api_key: str, module: str ) -> bool: """Check whether an API key can access the requested module.""" - key_hash = hashlib.sha256(api_key.encode()).hexdigest() + key_hash = hash_api_key(api_key) api_key_record = await self.repository.get_by_key_hash(session, key_hash) if not api_key_record or not api_key_record.is_valid(): return False - # Allow access if all modules are enabled or the specific module is present. enabled_modules = api_key_record.enabled_modules or [] return "all" in enabled_modules or module in enabled_modules @@ -329,7 +402,7 @@ async def toggle_api_key( await session.refresh(api_key) if not api_key.is_active: - await identity_cache.invalidate_apikey( + await api_key_identity_cache.invalidate_api_key( redis_pool_manager.get_redis_service(), user_id, api_key.key_hash, diff --git a/apps/api/app/services/billing/stripe_service.py b/apps/api/app/services/billing/stripe_service.py index 3d10b5688..be1792a2e 100644 --- a/apps/api/app/services/billing/stripe_service.py +++ b/apps/api/app/services/billing/stripe_service.py @@ -6,12 +6,11 @@ import stripe from app.repositories.payment_record_repository import PaymentRecordRepository from app.services.billing.price_config_service import PriceConfigService -from app.services.rate_limit.identity_cache import identity_cache from app.services.rate_limit.tier_service import TierService from sqlalchemy import func, select from sqlalchemy.ext.asyncio import AsyncSession -from shared.core.config import redis_pool_manager, settings +from shared.core.config import settings from shared.core.exceptions.domain_exceptions import ( AuthException, KnowhereException, @@ -385,11 +384,6 @@ async def _handle_checkout_completed( await db.commit() await db.refresh(payment_record) - await identity_cache.invalidate_user( - redis_pool_manager.get_redis_service(), - user_id, - ) - logger.info( f"Credits pack purchase succeeded: user_id={user_id}, credits={credits_amount}, price_id={price_id}" ) @@ -527,11 +521,6 @@ async def _handle_payment_intent_succeeded( await db.commit() await db.refresh(payment_record) - await identity_cache.invalidate_user( - redis_pool_manager.get_redis_service(), - user_id, - ) - logger.info( f"buy credits success: user_id={user_id}, credits={credits_amount}, payment_intent_id={payment_intent_id}" ) diff --git a/apps/api/app/services/guest/guest_registration_service.py b/apps/api/app/services/guest/guest_registration_service.py index 351663643..6cd16ee63 100644 --- a/apps/api/app/services/guest/guest_registration_service.py +++ b/apps/api/app/services/guest/guest_registration_service.py @@ -25,6 +25,7 @@ GuestRegisterResponse, ) from shared.services.billing.credits_service import CreditsService +from shared.utils.api_keys import generate_api_key, hash_api_key, mask_api_key _GUEST_TIER: str = "guest" _GUEST_KEY_NAME_PREFIX: str = "guest-device" @@ -153,14 +154,11 @@ async def _create_api_key_without_commit( This avoids the internal commit inside APIKeyService.create_api_key() which would make the key durable before the device row is inserted. """ - import hashlib - import uuid - from shared.models.database.api_key import APIKey - api_key = f"sk_{str(uuid.uuid4()).replace('-', '')}" - key_hash = hashlib.sha256(api_key.encode()).hexdigest() - key_mask = self._api_key_service._mask_api_key(api_key) + api_key = generate_api_key() + key_hash = hash_api_key(api_key) + key_mask = mask_api_key(api_key) api_key_record = APIKey( user_id=user_id, @@ -264,13 +262,11 @@ def _raise_existing_device_conflict(cls, device_id: str) -> NoReturn: @staticmethod async def _resolve_api_key_id(session: AsyncSession, api_key: str) -> str | None: """Resolve the DB id for a just-created API key by its hash.""" - import hashlib - from sqlalchemy import select from shared.models.database.api_key import APIKey - key_hash = hashlib.sha256(api_key.encode()).hexdigest() + key_hash = hash_api_key(api_key) result = await session.execute( select(APIKey.id).where(APIKey.key_hash == key_hash).limit(1) ) diff --git a/apps/api/app/services/rate_limit/dependencies.py b/apps/api/app/services/rate_limit/dependencies.py index c34b1c10b..cd2a29e6d 100644 --- a/apps/api/app/services/rate_limit/dependencies.py +++ b/apps/api/app/services/rate_limit/dependencies.py @@ -5,8 +5,8 @@ require_billing_limits -> with_current_user -> get_current_user_id -> get_db -``with_current_user`` resolves identity (user_id + user_tier), caches it in -Redis, and enforces the matched system limit (Layer 0). +``with_current_user`` resolves the user's billing tier, caches it in Redis, +and enforces the matched system limit (Layer 0). ``require_billing_limits`` enforces billing RPM (Layer 1) when billing is enabled and yields control to the route handler. Concurrency (Layer 2) and @@ -14,9 +14,7 @@ only when billing is enabled. """ -import hashlib import math -from datetime import datetime, timezone from typing import AsyncGenerator from app.core.dependencies import get_current_user_id @@ -41,7 +39,6 @@ ) from shared.core.logging import log_context from shared.core.state_machine.states import JobStatus -from shared.models.database.api_key import APIKey from shared.models.database.job import Job from shared.models.database.user_balance import UserBalance @@ -83,16 +80,6 @@ async def _resolve_user_tier_from_db(user_id: str) -> str: return _DEFAULT_TIER -def _extract_bearer_token(authorization: str | None) -> str | None: - """Extract bearer token from Authorization header.""" - if not authorization: - return None - scheme, _, token = authorization.partition(" ") - if scheme.lower() != "bearer" or not token: - return None - return token - - def _get_route_path(request: Request) -> str: """Return the request path without the application's root_path prefix.""" scope_path: str = request.scope.get("path", request.url.path) @@ -116,30 +103,6 @@ def _get_route_limit_identifier(request: Request) -> str: return _get_route_path(request) -async def _resolve_apikey_cache_ttl_seconds(api_key_hash: str) -> int: - """Resolve cache TTL for API key identity (max 1 hour).""" - max_ttl_seconds = 3600 - try: - async with get_db_context() as db: - result = await db.execute( - select(APIKey.expires_at) - .where(APIKey.key_hash == api_key_hash) - .limit(1) - ) - expires_at = result.scalar_one_or_none() - if expires_at is None: - return max_ttl_seconds - - # APIKey.expires_at is stored as UTC-naive datetime. - if expires_at.tzinfo is None: - expires_at = expires_at.replace(tzinfo=timezone.utc) - now = datetime.now(timezone.utc) - remaining = int((expires_at - now).total_seconds()) - return max(1, min(max_ttl_seconds, remaining)) - except Exception: - return max_ttl_seconds - - # --------------------------------------------------------------------------- # with_current_user -- Layer 0 (matched system limit) # --------------------------------------------------------------------------- @@ -154,7 +117,7 @@ async def with_current_user( Steps: 1. ``get_current_user_id`` already authenticated the user (401 on failure). - 2. Resolve ``user_tier`` from the identity cache; fall back to DB on + 2. Resolve ``user_tier`` from the tier cache; fall back to DB on cache miss or Redis error. 3. If ``RATE_LIMIT_ENABLED=false`` is set, return immediately. 4. Check the matched system limit via the rate limiter (fail-open on @@ -163,74 +126,23 @@ async def with_current_user( redis_service = redis_pool_manager.get_redis_service() # -- Resolve user_tier (cache -> DB fallback) -- - user_tier: str | None = getattr(request.state, "cached_user_tier", None) - stashed_user_id: str | None = getattr(request.state, "user_id", None) - cached_identity_hit: bool | None = getattr( - request.state, "cached_identity_hit", None - ) - - if isinstance(user_tier, str) and stashed_user_id == user_id: - if cached_identity_hit is False: - token = _extract_bearer_token(request.headers.get("authorization")) - api_key_hash = None - is_api_key_auth = isinstance(token, str) and token.startswith("sk_") - if token is not None and is_api_key_auth: - api_key_hash = hashlib.sha256(token.encode()).hexdigest() - if is_api_key_auth and api_key_hash: - try: - ttl_seconds = await _resolve_apikey_cache_ttl_seconds(api_key_hash) - await identity_cache.set_apikey_identity( - redis_service, - api_key_hash, - user_id, - user_tier, - ttl_seconds=ttl_seconds, - ) - except Exception: - logger.warning( - "rate_limit: failed to backfill API key identity cache " - "for user_id={}", - user_id, - ) - else: - token = _extract_bearer_token(request.headers.get("authorization")) - api_key_hash = None - is_api_key_auth = isinstance(token, str) and token.startswith("sk_") - if token is not None and is_api_key_auth: - api_key_hash = hashlib.sha256(token.encode()).hexdigest() - cache_key: str = ( - identity_cache._apikey_key(api_key_hash) - if is_api_key_auth and api_key_hash - else identity_cache._jwt_key(user_id) + user_tier: str | None = None + try: + cached: dict[str, str] | None = await identity_cache.get_user_tier( + redis_service, user_id ) - try: - cached: dict | None = await identity_cache.get_cached_identity( - redis_service, cache_key - ) - if cached is not None: - user_tier = cached.get("user_tier", _DEFAULT_TIER) - else: - user_tier = await _resolve_user_tier_from_db(user_id) - if is_api_key_auth and api_key_hash: - ttl_seconds = await _resolve_apikey_cache_ttl_seconds(api_key_hash) - await identity_cache.set_apikey_identity( - redis_service, - api_key_hash, - user_id, - user_tier, - ttl_seconds=ttl_seconds, - ) - else: - await identity_cache.set_jwt_identity( - redis_service, user_id, user_tier - ) - except Exception: - logger.warning( - "rate_limit: Redis error during identity resolution, " - "falling back to DB for user_id={}", - user_id, - ) + if cached is not None: + user_tier = cached.get("user_tier", _DEFAULT_TIER) + else: user_tier = await _resolve_user_tier_from_db(user_id) + await identity_cache.set_user_tier(redis_service, user_id, user_tier) + except Exception: + logger.warning( + "rate_limit: Redis error during tier resolution, " + "falling back to DB for user_id={}", + user_id, + ) + user_tier = await _resolve_user_tier_from_db(user_id) if user_tier is None: user_tier = _DEFAULT_TIER diff --git a/apps/api/app/services/rate_limit/identity_cache.py b/apps/api/app/services/rate_limit/identity_cache.py index 9e6ffae01..5ec04696e 100644 --- a/apps/api/app/services/rate_limit/identity_cache.py +++ b/apps/api/app/services/rate_limit/identity_cache.py @@ -1,69 +1,35 @@ -""" -Redis-backed identity cache for user_id + user_tier resolution. +"""Redis-backed user-tier cache for rate-limit identity resolution.""" -Caches the mapping from authentication credentials (JWT user_id or API key hash) -to the resolved identity (user_id, user_tier) so that -tier lookups do not hit the database on every request. +from __future__ import annotations -Key patterns (all prefixed with REDIS_KEY_PREFIX from config): - JWT: {REDIS_KEY_PREFIX}identity:jwt:{user_id} - API key: {REDIS_KEY_PREFIX}identity:apikey:{api_key_hash} - Reverse: {REDIS_KEY_PREFIX}identity:apikeys:{user_id} -""" +from typing import TYPE_CHECKING -import json -from typing import Optional - -from app.services.rate_limit.config import REDIS_KEY_PREFIX from loguru import logger -from shared.services.redis.redis_service import RedisService - -# Default TTL for JWT identity cache entries (1 hour). -_JWT_TTL_SECONDS: int = 3600 +if TYPE_CHECKING: + from shared.services.redis.redis_service import RedisService -# Upper bound TTL for API-key identity cache entries (1 hour). -_APIKEY_MAX_TTL_SECONDS: int = 3600 +_USER_TIER_TTL_SECONDS: int = 3600 class IdentityCache: - """Redis-backed identity cache for user_id + user_tier resolution.""" - - # ------------------------------------------------------------------ - # Key builders - # ------------------------------------------------------------------ - - @staticmethod - def _jwt_key(user_id: str) -> str: - return f"{REDIS_KEY_PREFIX}identity:jwt:{user_id}" + """Cache resolved rate-limit user tier by user_id.""" @staticmethod - def _apikey_key(api_key_hash: str) -> str: - return f"{REDIS_KEY_PREFIX}identity:apikey:{api_key_hash}" + def get_user_tier_key(user_id: str) -> str: + """Return the Redis key for a user's rate-limit tier.""" + return f"identity:user-tier:{user_id}" - @staticmethod - def _reverse_key(user_id: str) -> str: - return f"{REDIS_KEY_PREFIX}identity:apikeys:{user_id}" - - # ------------------------------------------------------------------ - # Read - # ------------------------------------------------------------------ - - async def get_cached_identity( + async def get_user_tier( self, redis: RedisService, - cache_key: str, - ) -> Optional[dict]: + user_id: str, + ) -> dict[str, str] | None: """Return cached ``{user_id, user_tier}`` or ``None`` on miss.""" + cache_key: str = self.get_user_tier_key(user_id) try: - raw: Optional[str] = await redis.get(cache_key) - if raw is None: - return None - # RedisService.get already attempts JSON parse, but the - # value may come back as a dict directly. - if isinstance(raw, dict): - return raw - return json.loads(raw) + raw_identity: object = await redis.get(cache_key) + return self._coerce_identity(raw_identity) except Exception: logger.warning( "identity_cache: failed to read cache_key={}", @@ -71,121 +37,48 @@ async def get_cached_identity( ) return None - # ------------------------------------------------------------------ - # Write -- JWT - # ------------------------------------------------------------------ - - async def set_jwt_identity( - self, - redis: RedisService, - user_id: str, - user_tier: str, - ) -> None: - """Cache identity for a JWT-authenticated user (1 hr TTL).""" - key: str = self._jwt_key(user_id) - payload: dict = {"user_id": user_id, "user_tier": user_tier} - try: - await redis.set(key, payload, ttl=_JWT_TTL_SECONDS) - except Exception: - logger.warning( - "identity_cache: failed to set jwt identity user_id={}", - user_id, - ) - - # ------------------------------------------------------------------ - # Write -- API key - # ------------------------------------------------------------------ - - async def set_apikey_identity( + async def set_user_tier( self, redis: RedisService, - api_key_hash: str, user_id: str, user_tier: str, - ttl_seconds: int, ) -> None: - """Cache identity for an API-key-authenticated user. - - TTL is ``min(APIKEY_MAX_TTL, api_key_remaining_ttl)`` so the - cache never outlives the key itself. Also maintains a reverse - index (SET) of all cached API-key hashes per user for bulk - invalidation. - """ - effective_ttl: int = min(_APIKEY_MAX_TTL_SECONDS, ttl_seconds) - key: str = self._apikey_key(api_key_hash) - payload: dict = {"user_id": user_id, "user_tier": user_tier} + """Cache rate-limit tier for a user.""" + key: str = self.get_user_tier_key(user_id) + payload: dict[str, str] = {"user_id": user_id, "user_tier": user_tier} try: - await redis.set(key, payload, ttl=effective_ttl) - # Maintain reverse index so invalidate_user can find all - # API-key cache entries belonging to this user. - reverse_key: str = self._reverse_key(user_id) - await redis.sadd(reverse_key, api_key_hash) - # Keep reverse index TTL at least as long as the longest - # surviving API-key cache entry for this user. - current_ttl = await redis.ttl(reverse_key) - if current_ttl in (-2, -1) or current_ttl < effective_ttl: - await redis.expire(reverse_key, effective_ttl) + await redis.set(key, payload, ttl=_USER_TIER_TTL_SECONDS) except Exception: logger.warning( - "identity_cache: failed to set apikey identity " - "api_key_hash={}, user_id={}", - api_key_hash, + "identity_cache: failed to set user tier for user_id={}", user_id, ) - # ------------------------------------------------------------------ - # Invalidation - # ------------------------------------------------------------------ - async def invalidate_user( self, redis: RedisService, user_id: str, ) -> None: - """Full invalidation: JWT cache + all API-key caches + reverse index.""" + """Delete cached rate-limit tier for a user.""" try: - # 1. Delete JWT cache - jwt_key: str = self._jwt_key(user_id) - await redis.delete(jwt_key) - - # 2. Collect all cached API-key hashes from reverse index - reverse_key: str = self._reverse_key(user_id) - api_key_hashes: set = await redis.smembers(reverse_key) - - # 3. Delete each API-key cache entry - for api_key_hash in api_key_hashes: - apikey_key: str = self._apikey_key(str(api_key_hash)) - await redis.delete(apikey_key) - - # 4. Delete the reverse index itself - await redis.delete(reverse_key) + await redis.delete(self.get_user_tier_key(user_id)) except Exception: logger.warning( "identity_cache: failed to invalidate user_id={}", user_id, ) - async def invalidate_apikey( - self, - redis: RedisService, - user_id: str, - api_key_hash: str, - ) -> None: - """Delete a single API-key cache entry and remove from reverse index.""" - try: - apikey_key: str = self._apikey_key(api_key_hash) - await redis.delete(apikey_key) + def _coerce_identity(self, raw_identity: object) -> dict[str, str] | None: + """Return a typed identity from current or legacy Redis values.""" + if not isinstance(raw_identity, dict): + return None - reverse_key: str = self._reverse_key(user_id) - await redis.srem(reverse_key, api_key_hash) - except Exception: - logger.warning( - "identity_cache: failed to invalidate apikey " - "api_key_hash={}, user_id={}", - api_key_hash, - user_id, - ) + raw_user_id: object = raw_identity.get("user_id") + raw_user_tier: object = raw_identity.get("user_tier") + if not isinstance(raw_user_id, str) or not isinstance(raw_user_tier, str): + return None + + return {"user_id": raw_user_id, "user_tier": raw_user_tier} -# Module-level singleton so callers can import directly. identity_cache = IdentityCache() diff --git a/apps/api/app/services/rate_limit/tier_service.py b/apps/api/app/services/rate_limit/tier_service.py index a9d73a13e..67fe20408 100644 --- a/apps/api/app/services/rate_limit/tier_service.py +++ b/apps/api/app/services/rate_limit/tier_service.py @@ -7,10 +7,12 @@ from app.services.rate_limit.config import RateLimitConfig from app.services.rate_limit.data_structures import TierLimits +from app.services.rate_limit.identity_cache import identity_cache from loguru import logger from sqlalchemy import func, select, update from sqlalchemy.ext.asyncio import AsyncSession +from shared.core.config import redis_pool_manager from shared.models.database.payment_record import PaymentRecord from shared.models.database.tier_limit import TierLimit from shared.models.database.user_balance import UserBalance @@ -61,6 +63,7 @@ async def refresh_tier(user_id: str, session: AsyncSession) -> str: .values(user_tier=new_tier) ) await session.execute(stmt_update) + await TierService._invalidate_user_tier_cache(user_id) logger.info( "Tier refreshed: user_id=%s total_micro=%d new_tier=%s", @@ -78,3 +81,17 @@ def get_limits(user_tier: str) -> Optional[TierLimits]: """ config = RateLimitConfig.get_instance() return config.tier_map.get(user_tier) + + @staticmethod + async def _invalidate_user_tier_cache(user_id: str) -> None: + """Invalidate cached tier data without touching API-key auth cache.""" + try: + await identity_cache.invalidate_user( + redis_pool_manager.get_redis_service(), + user_id, + ) + except Exception: + logger.warning( + "Tier refresh cache invalidation failed for user_id={}", + user_id, + ) diff --git a/apps/api/tests/contract/test_identity_cache_contract.py b/apps/api/tests/contract/test_identity_cache_contract.py new file mode 100644 index 000000000..82abecf86 --- /dev/null +++ b/apps/api/tests/contract/test_identity_cache_contract.py @@ -0,0 +1,117 @@ +import builtins +from typing import TYPE_CHECKING, cast + +import pytest + +from app.services.auth.api_key_identity_cache import APIKeyIdentityCache +from app.services.rate_limit.identity_cache import IdentityCache + +if TYPE_CHECKING: + from shared.services.redis.redis_service import RedisService +else: + RedisService = object + + +class FakeRedisService: + def __init__(self) -> None: + self.values: dict[str, object] = {} + self.sets: dict[str, set[str]] = {} + self.ttls: dict[str, int] = {} + + async def get(self, key: str) -> object | None: + return self.values.get(key) + + async def set(self, key: str, value: object, ttl: int | None = None) -> bool: + self.values[key] = value + if ttl is not None: + self.ttls[key] = ttl + return True + + async def delete(self, *keys: str) -> int: + deleted_count: int = 0 + for key in keys: + deleted_value: object | None = self.values.pop(key, None) + deleted_set: set[str] | None = self.sets.pop(key, None) + self.ttls.pop(key, None) + if deleted_value is not None or deleted_set is not None: + deleted_count += 1 + return deleted_count + + async def sadd(self, key: str, *values: object) -> int: + members: builtins.set[str] = self.sets.setdefault(key, set()) + previous_size: int = len(members) + members.update(str(value) for value in values) + return len(members) - previous_size + + async def srem(self, key: str, *values: object) -> int: + members: builtins.set[str] = self.sets.setdefault(key, set()) + removed_count: int = 0 + for value in values: + if str(value) in members: + members.remove(str(value)) + removed_count += 1 + return removed_count + + async def smembers(self, key: str) -> builtins.set[str]: + return set(self.sets.get(key, set())) + + async def ttl(self, key: str) -> int: + return self.ttls.get(key, -2) + + async def expire(self, key: str, ttl: int) -> bool: + self.ttls[key] = ttl + return True + + +@pytest.mark.asyncio +async def test_api_key_identity_cache_should_store_user_id_without_tier() -> None: + cache = APIKeyIdentityCache() + fake_redis = FakeRedisService() + redis = cast(RedisService, fake_redis) + + await cache.set_user_id( + redis, + api_key_hash="hash-one", + user_id="user-one", + ttl_seconds=7200, + ) + + assert await cache.get_user_id(redis, "hash-one") == "user-one" + assert fake_redis.values[cache.get_cache_key("hash-one")] == "user-one" + assert fake_redis.ttls[cache.get_cache_key("hash-one")] == 3600 + assert fake_redis.sets[cache.get_reverse_key("user-one")] == {"hash-one"} + + +@pytest.mark.asyncio +async def test_api_key_identity_cache_invalidation_should_not_touch_tier_cache() -> None: + api_key_cache = APIKeyIdentityCache() + tier_cache = IdentityCache() + fake_redis = FakeRedisService() + redis = cast(RedisService, fake_redis) + + await api_key_cache.set_user_id(redis, "hash-one", "user-one", ttl_seconds=300) + await tier_cache.set_user_tier(redis, "user-one", "tier_5") + + await api_key_cache.invalidate_api_key(redis, "user-one", "hash-one") + + assert await api_key_cache.get_user_id(redis, "hash-one") is None + assert await tier_cache.get_user_tier(redis, "user-one") == { + "user_id": "user-one", + "user_tier": "tier_5", + } + + +@pytest.mark.asyncio +async def test_user_tier_cache_invalidation_should_not_touch_api_key_cache() -> None: + api_key_cache = APIKeyIdentityCache() + tier_cache = IdentityCache() + fake_redis = FakeRedisService() + redis = cast(RedisService, fake_redis) + + await api_key_cache.set_user_id(redis, "hash-one", "user-one", ttl_seconds=300) + await tier_cache.set_user_tier(redis, "user-one", "tier_5") + + await tier_cache.invalidate_user(redis, "user-one") + + assert await api_key_cache.get_user_id(redis, "hash-one") == "user-one" + assert await tier_cache.get_user_tier(redis, "user-one") is None diff --git a/packages/shared-python/shared/tests/utils/test_api_keys.py b/packages/shared-python/shared/tests/utils/test_api_keys.py new file mode 100644 index 000000000..c51d924f1 --- /dev/null +++ b/packages/shared-python/shared/tests/utils/test_api_keys.py @@ -0,0 +1,34 @@ +from shared.utils.api_keys import ( + API_KEY_PREFIX, + generate_api_key, + hash_api_key, + is_api_key_token, + mask_api_key, +) + + +def test_generate_api_key_should_use_api_key_prefix_and_random_secret() -> None: + first_api_key: str = generate_api_key() + second_api_key: str = generate_api_key() + + assert first_api_key.startswith(API_KEY_PREFIX) + assert second_api_key.startswith(API_KEY_PREFIX) + assert first_api_key != second_api_key + assert len(first_api_key) > len(API_KEY_PREFIX) + 32 + + +def test_hash_api_key_should_return_deterministic_sha256_lookup_hash() -> None: + api_key: str = "sk_contract_test_secret" + + assert hash_api_key(api_key) == hash_api_key(api_key) + assert len(hash_api_key(api_key)) == 64 + + +def test_mask_api_key_should_hide_middle_characters() -> None: + assert mask_api_key("sk_1234567890abcdef") == "sk_12345•••••••cdef" + + +def test_is_api_key_token_should_match_only_api_key_prefix() -> None: + assert is_api_key_token("sk_test") is True + assert is_api_key_token("jwt_test") is False + assert is_api_key_token(None) is False diff --git a/packages/shared-python/shared/utils/api_keys.py b/packages/shared-python/shared/utils/api_keys.py new file mode 100644 index 000000000..c7feec616 --- /dev/null +++ b/packages/shared-python/shared/utils/api_keys.py @@ -0,0 +1,30 @@ +"""API key generation, masking, and hashing helpers.""" + +from hashlib import sha256 +from secrets import token_urlsafe +from typing import TypeGuard + +API_KEY_PREFIX: str = "sk_" +API_KEY_RANDOM_BYTES: int = 32 + + +def generate_api_key() -> str: + """Generate a new plaintext API key with cryptographic randomness.""" + return f"{API_KEY_PREFIX}{token_urlsafe(API_KEY_RANDOM_BYTES)}" + + +def hash_api_key(api_key: str) -> str: + """Return a deterministic SHA-256 digest for API key lookup.""" + return sha256(api_key.encode("utf-8")).hexdigest() + + +def mask_api_key(api_key: str) -> str: + """Mask an API key, exposing only the first 8 and last 4 characters.""" + if len(api_key) < 12: + return api_key + return api_key[:8] + "•" * (len(api_key) - 12) + api_key[-4:] + + +def is_api_key_token(token: object) -> TypeGuard[str]: + """Return whether a bearer token has the API-key prefix.""" + return isinstance(token, str) and token.startswith(API_KEY_PREFIX) From d69a45388c278cb286926761af73454a4b1ef4a4 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 16:41:03 +0800 Subject: [PATCH 18/32] refactor: remove unused health check endpoints and related tests --- apps/api/app/api/v1/api_v1.py | 4 - apps/api/app/api/v1/health.py | 7 -- apps/api/app/api/v1/routes/api_key.py | 24 ---- .../services/auth/api_key_identity_cache.py | 18 --- apps/api/app/services/auth/api_key_service.py | 63 +--------- .../tests/contract/test_api_key_contract.py | 51 -------- .../contract/test_identity_cache_contract.py | 117 ------------------ .../shared-python/shared/utils/api_keys.py | 2 +- 8 files changed, 2 insertions(+), 284 deletions(-) delete mode 100644 apps/api/app/api/v1/health.py delete mode 100644 apps/api/tests/contract/test_identity_cache_contract.py diff --git a/apps/api/app/api/v1/api_v1.py b/apps/api/app/api/v1/api_v1.py index b842e3900..023fb7621 100644 --- a/apps/api/app/api/v1/api_v1.py +++ b/apps/api/app/api/v1/api_v1.py @@ -2,7 +2,6 @@ API v1 route registry. """ -from app.api.v1 import health from app.api.v1.routes import ( api_key, documents, @@ -59,9 +58,6 @@ qstash_callbacks.router, prefix="/webhooks", tags=["QStash Callbacks"] ) -# Health check -api_router.include_router(health.router, prefix="/health", tags=["Health"]) - # Version info api_router.include_router(version.router, tags=["Version"]) diff --git a/apps/api/app/api/v1/health.py b/apps/api/app/api/v1/health.py deleted file mode 100644 index d42aed21e..000000000 --- a/apps/api/app/api/v1/health.py +++ /dev/null @@ -1,7 +0,0 @@ -""" -Database health API endpoints. -""" - -from fastapi import APIRouter - -router = APIRouter() diff --git a/apps/api/app/api/v1/routes/api_key.py b/apps/api/app/api/v1/routes/api_key.py index 7564d189c..b39ca1366 100644 --- a/apps/api/app/api/v1/routes/api_key.py +++ b/apps/api/app/api/v1/routes/api_key.py @@ -98,30 +98,6 @@ async def list_api_keys( ) -@router.post("/regenerate", summary="Regenerate an API key") -async def regenerate_api_key( - request: RegenerateAPIKeyRequest, - current_user: CurrentUser = Depends(with_current_user), - db: AsyncSession = Depends(get_db), -): - """Regenerate an API key.""" - api_key_service = APIKeyService() - - try: - new_api_key = await api_key_service.regenerate_api_key( - session=db, api_key_id=request.api_key_id, user_id=current_user.user_id - ) - - return {"api_key": new_api_key, "message": "API key regenerated"} - - except NotFoundException: - raise - except Exception as e: - raise APIKeyOperationException( - internal_message=f"Failed to regenerate API Key: {str(e)}" - ) - - @router.post("/revoke", summary="Revoke an API key") async def revoke_api_key( request: RevokeAPIKeyRequest, diff --git a/apps/api/app/services/auth/api_key_identity_cache.py b/apps/api/app/services/auth/api_key_identity_cache.py index 95308853c..683113abc 100644 --- a/apps/api/app/services/auth/api_key_identity_cache.py +++ b/apps/api/app/services/auth/api_key_identity_cache.py @@ -82,24 +82,6 @@ async def invalidate_api_key( user_id, ) - async def invalidate_user( - self, - redis: RedisService, - user_id: str, - ) -> None: - """Delete all API-key identity cache entries for a user.""" - try: - reverse_key: str = self.get_reverse_key(user_id) - api_key_hashes: set[object] = await redis.smembers(reverse_key) - for api_key_hash in api_key_hashes: - await redis.delete(self.get_cache_key(str(api_key_hash))) - await redis.delete(reverse_key) - except Exception: - logger.warning( - "api_key_identity_cache: failed to invalidate user_id={}", - user_id, - ) - def _coerce_user_id(self, raw_user_id: object) -> str | None: """Return a typed user ID from current or legacy Redis values.""" if isinstance(raw_user_id, str): diff --git a/apps/api/app/services/auth/api_key_service.py b/apps/api/app/services/auth/api_key_service.py index b79baa3ae..c810eb7ca 100644 --- a/apps/api/app/services/auth/api_key_service.py +++ b/apps/api/app/services/auth/api_key_service.py @@ -99,21 +99,6 @@ async def create_api_key( return api_key - async def validate_api_key( - self, session: AsyncSession, api_key: str - ) -> Optional[str]: - """Validate API key against DB, return user_id or None.""" - identity = await self.get_identity(session, api_key) - return identity.user_id if identity is not None else None - - async def validate_api_key_identity( - self, - session: AsyncSession, - api_key: str, - ) -> Optional[APIKeyIdentity]: - """Validate API key and return the authenticated identity.""" - return await self.get_identity(session, api_key) - async def get_identity( self, session: AsyncSession, @@ -266,6 +251,7 @@ async def revoke_api_key( return success + # TODO, invalidate should not be best-effort async def _invalidate_revoked_api_key_cache_best_effort( self, user_id: str, @@ -303,53 +289,6 @@ async def list_user_api_keys( for api_key in api_keys ] - async def regenerate_api_key( - self, session: AsyncSession, api_key_id: str, user_id: str - ) -> str: - """Regenerate an API key.""" - api_key = await self.repository.get_by_id(session, api_key_id) - if not api_key or api_key.user_id != user_id: - raise NotFoundException( - resource="APIKey", - resource_id=api_key_id, - internal_message="API Key not found or does not belong to user", - ) - - new_api_key = generate_api_key() - new_key_hash = hash_api_key(new_api_key) - new_key_mask = mask_api_key(new_api_key) - - await session.execute( - update(APIKey) - .where(APIKey.id == api_key_id) - .values( - key_hash=new_key_hash, - key_mask=new_key_mask, - ) - ) - await session.commit() - - await api_key_identity_cache.invalidate_api_key( - redis_pool_manager.get_redis_service(), - user_id, - api_key.key_hash, - ) - - return new_api_key - - async def check_module_permission( - self, session: AsyncSession, api_key: str, module: str - ) -> bool: - """Check whether an API key can access the requested module.""" - key_hash = hash_api_key(api_key) - api_key_record = await self.repository.get_by_key_hash(session, key_hash) - - if not api_key_record or not api_key_record.is_valid(): - return False - - enabled_modules = api_key_record.enabled_modules or [] - return "all" in enabled_modules or module in enabled_modules - def _schedule_last_used_update(self, api_key_id: str) -> None: """Schedule a best-effort background update for api_keys.last_used_at.""" try: diff --git a/apps/api/tests/contract/test_api_key_contract.py b/apps/api/tests/contract/test_api_key_contract.py index 21b693fb3..eef3053ba 100644 --- a/apps/api/tests/contract/test_api_key_contract.py +++ b/apps/api/tests/contract/test_api_key_contract.py @@ -77,57 +77,6 @@ async def test_should_revoke_a_created_api_key_through_http_only( assert error["message"] == "Invalid API Key" assert "details" not in error - -@pytest.mark.asyncio -async def test_should_regenerate_an_api_key_and_invalidate_the_previous_raw_key( - developer_api_client_factory: Callable[ - [], AbstractAsyncContextManager[AsyncClient] - ], -) -> None: - create_payload: dict[str, object] = { - "name": f"contract-regenerate-{uuid4().hex[:8]}", - "enabled_modules": ["jobs"], - } - - async with developer_api_client_factory() as api_client: - create_response = await api_client.post("/api/v1/auth/create", json=create_payload) - assert create_response.status_code == 200 - create_response_json = cast(dict[str, object], create_response.json()) - old_api_key = cast(str, create_response_json["api_key"]) - - list_response = await api_client.get("/api/v1/auth/list") - assert list_response.status_code == 200 - list_response_json = cast(dict[str, object], list_response.json()) - api_keys = cast(list[dict[str, object]], list_response_json["api_keys"]) - created_api_key = next( - api_key - for api_key in api_keys - if api_key["name"] == create_payload["name"] - ) - created_api_key_id = cast(str, created_api_key["id"]) - - regenerate_response = await api_client.post( - "/api/v1/auth/regenerate", - json={"api_key_id": created_api_key_id}, - ) - assert regenerate_response.status_code == 200 - regenerate_response_json = cast(dict[str, object], regenerate_response.json()) - new_api_key = cast(str, regenerate_response_json["api_key"]) - - api_client.headers.update({"Authorization": f"Bearer {old_api_key}"}) - old_key_response = await api_client.get("/api/v1/jobs") - - api_client.headers.update({"Authorization": f"Bearer {new_api_key}"}) - new_key_response = await api_client.get("/api/v1/jobs") - - assert regenerate_response_json["message"] == "API key regenerated" - assert new_api_key.startswith("sk_") - assert new_api_key != old_api_key - assert old_key_response.status_code == 401 - assert old_key_response.json()["error"]["code"] == "UNAUTHENTICATED" - assert new_key_response.status_code == 200 - - @pytest.mark.asyncio async def test_should_return_owned_api_key_metadata( developer_api_client_factory: Callable[ diff --git a/apps/api/tests/contract/test_identity_cache_contract.py b/apps/api/tests/contract/test_identity_cache_contract.py deleted file mode 100644 index 82abecf86..000000000 --- a/apps/api/tests/contract/test_identity_cache_contract.py +++ /dev/null @@ -1,117 +0,0 @@ -import builtins -from typing import TYPE_CHECKING, cast - -import pytest - -from app.services.auth.api_key_identity_cache import APIKeyIdentityCache -from app.services.rate_limit.identity_cache import IdentityCache - -if TYPE_CHECKING: - from shared.services.redis.redis_service import RedisService -else: - RedisService = object - - -class FakeRedisService: - def __init__(self) -> None: - self.values: dict[str, object] = {} - self.sets: dict[str, set[str]] = {} - self.ttls: dict[str, int] = {} - - async def get(self, key: str) -> object | None: - return self.values.get(key) - - async def set(self, key: str, value: object, ttl: int | None = None) -> bool: - self.values[key] = value - if ttl is not None: - self.ttls[key] = ttl - return True - - async def delete(self, *keys: str) -> int: - deleted_count: int = 0 - for key in keys: - deleted_value: object | None = self.values.pop(key, None) - deleted_set: set[str] | None = self.sets.pop(key, None) - self.ttls.pop(key, None) - if deleted_value is not None or deleted_set is not None: - deleted_count += 1 - return deleted_count - - async def sadd(self, key: str, *values: object) -> int: - members: builtins.set[str] = self.sets.setdefault(key, set()) - previous_size: int = len(members) - members.update(str(value) for value in values) - return len(members) - previous_size - - async def srem(self, key: str, *values: object) -> int: - members: builtins.set[str] = self.sets.setdefault(key, set()) - removed_count: int = 0 - for value in values: - if str(value) in members: - members.remove(str(value)) - removed_count += 1 - return removed_count - - async def smembers(self, key: str) -> builtins.set[str]: - return set(self.sets.get(key, set())) - - async def ttl(self, key: str) -> int: - return self.ttls.get(key, -2) - - async def expire(self, key: str, ttl: int) -> bool: - self.ttls[key] = ttl - return True - - -@pytest.mark.asyncio -async def test_api_key_identity_cache_should_store_user_id_without_tier() -> None: - cache = APIKeyIdentityCache() - fake_redis = FakeRedisService() - redis = cast(RedisService, fake_redis) - - await cache.set_user_id( - redis, - api_key_hash="hash-one", - user_id="user-one", - ttl_seconds=7200, - ) - - assert await cache.get_user_id(redis, "hash-one") == "user-one" - assert fake_redis.values[cache.get_cache_key("hash-one")] == "user-one" - assert fake_redis.ttls[cache.get_cache_key("hash-one")] == 3600 - assert fake_redis.sets[cache.get_reverse_key("user-one")] == {"hash-one"} - - -@pytest.mark.asyncio -async def test_api_key_identity_cache_invalidation_should_not_touch_tier_cache() -> None: - api_key_cache = APIKeyIdentityCache() - tier_cache = IdentityCache() - fake_redis = FakeRedisService() - redis = cast(RedisService, fake_redis) - - await api_key_cache.set_user_id(redis, "hash-one", "user-one", ttl_seconds=300) - await tier_cache.set_user_tier(redis, "user-one", "tier_5") - - await api_key_cache.invalidate_api_key(redis, "user-one", "hash-one") - - assert await api_key_cache.get_user_id(redis, "hash-one") is None - assert await tier_cache.get_user_tier(redis, "user-one") == { - "user_id": "user-one", - "user_tier": "tier_5", - } - - -@pytest.mark.asyncio -async def test_user_tier_cache_invalidation_should_not_touch_api_key_cache() -> None: - api_key_cache = APIKeyIdentityCache() - tier_cache = IdentityCache() - fake_redis = FakeRedisService() - redis = cast(RedisService, fake_redis) - - await api_key_cache.set_user_id(redis, "hash-one", "user-one", ttl_seconds=300) - await tier_cache.set_user_tier(redis, "user-one", "tier_5") - - await tier_cache.invalidate_user(redis, "user-one") - - assert await api_key_cache.get_user_id(redis, "hash-one") == "user-one" - assert await tier_cache.get_user_tier(redis, "user-one") is None diff --git a/packages/shared-python/shared/utils/api_keys.py b/packages/shared-python/shared/utils/api_keys.py index c7feec616..3d30898f2 100644 --- a/packages/shared-python/shared/utils/api_keys.py +++ b/packages/shared-python/shared/utils/api_keys.py @@ -7,7 +7,7 @@ API_KEY_PREFIX: str = "sk_" API_KEY_RANDOM_BYTES: int = 32 - +# TODO, use an alphanumeric api key def generate_api_key() -> str: """Generate a new plaintext API key with cryptographic randomness.""" return f"{API_KEY_PREFIX}{token_urlsafe(API_KEY_RANDOM_BYTES)}" From 1478552d1b5effa0a84e22010d619ac486a6362e Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 10:17:22 +0000 Subject: [PATCH 19/32] refactor: separate api key and tier caches --- apps/api/app/api/v1/routes/api_key.py | 11 +- apps/api/app/core/dependencies.py | 69 +---- .../services/auth/api_key_identity_cache.py | 106 -------- apps/api/app/services/auth/api_key_service.py | 237 +++++++++--------- .../guest/guest_registration_service.py | 2 - .../app/services/rate_limit/dependencies.py | 115 ++++----- .../app/services/rate_limit/identity_cache.py | 84 ------- .../app/services/rate_limit/tier_service.py | 97 ++++++- apps/api/tests/conftest.py | 3 + .../test_api_key_user_cache_contract.py | 115 +++++++++ .../contract/test_tier_service_contract.py | 79 ++++++ apps/api/tests/support/import_environment.py | 25 ++ 12 files changed, 499 insertions(+), 444 deletions(-) delete mode 100644 apps/api/app/services/auth/api_key_identity_cache.py delete mode 100644 apps/api/app/services/rate_limit/identity_cache.py create mode 100644 apps/api/tests/contract/test_api_key_user_cache_contract.py create mode 100644 apps/api/tests/contract/test_tier_service_contract.py create mode 100644 apps/api/tests/support/import_environment.py diff --git a/apps/api/app/api/v1/routes/api_key.py b/apps/api/app/api/v1/routes/api_key.py index b39ca1366..370521f7b 100644 --- a/apps/api/app/api/v1/routes/api_key.py +++ b/apps/api/app/api/v1/routes/api_key.py @@ -21,7 +21,6 @@ APIKeyResponse, CreateAPIKeyRequest, CreateAPIKeyResponse, - RegenerateAPIKeyRequest, RevokeAPIKeyRequest, ) @@ -35,7 +34,7 @@ async def create_api_key( db: AsyncSession = Depends(get_db), ): """Create an API key.""" - api_key_service = APIKeyService() + api_key_service = APIKeyService.get_instance() try: api_key = await api_key_service.create_api_key( @@ -69,7 +68,7 @@ async def list_api_keys( db: AsyncSession = Depends(get_db), ): """List API keys for the current user.""" - api_key_service = APIKeyService() + api_key_service = APIKeyService.get_instance() try: api_keys_data = await api_key_service.list_user_api_keys( @@ -105,7 +104,7 @@ async def revoke_api_key( db: AsyncSession = Depends(get_db), ): """Revoke an API key.""" - api_key_service = APIKeyService() + api_key_service = APIKeyService.get_instance() try: await api_key_service.revoke_api_key( @@ -130,7 +129,7 @@ async def get_api_key( db: AsyncSession = Depends(get_db), ): """Get details for a single API key.""" - api_key_service = APIKeyService() + api_key_service = APIKeyService.get_instance() try: api_key = await api_key_service.get_api_key( @@ -168,7 +167,7 @@ async def toggle_api_key( db: AsyncSession = Depends(get_db), ): """Enable or disable an API key.""" - api_key_service = APIKeyService() + api_key_service = APIKeyService.get_instance() try: success = await api_key_service.toggle_api_key( diff --git a/apps/api/app/core/dependencies.py b/apps/api/app/core/dependencies.py index 02038f5c9..0d03a4a95 100644 --- a/apps/api/app/core/dependencies.py +++ b/apps/api/app/core/dependencies.py @@ -1,6 +1,5 @@ import threading from datetime import timedelta -from fnmatch import fnmatch from typing import Any import jwt @@ -14,7 +13,6 @@ from shared.core.database import get_db from shared.core.exceptions.domain_exceptions import ( AuthException, - PermissionDeniedException, ) from shared.utils.api_keys import is_api_key_token @@ -27,22 +25,6 @@ # Cached PyJWKClient instance _jwks_client: PyJWKClient | None = None _jwks_client_lock = threading.Lock() -_GUEST_API_KEY_ALLOWED_ROUTE_PATTERNS: tuple[str, ...] = ( - "/v1/jobs", - "/v1/jobs/*", - "/v1/billing/credits", - "/v1/retrieval/query", - "/v1/documents", - "/v1/documents/*", - "/mcp", -) -_GUEST_API_KEY_REQUIRED_PERMISSION: str = ( - "jobs_documents_retrieval_mcp_or_billing_credits" -) -_GUEST_API_KEY_SCOPE_MESSAGE: str = ( - "Guest API keys can only access job, document, retrieval, MCP query, " - "and billing credits APIs" -) def _get_jwks_client() -> PyJWKClient: @@ -124,44 +106,6 @@ def decode_jwt_token(token: str) -> str: raise AuthException(user_message="Invalid token") -def _get_route_path(request: Request) -> str: - """Return the request path without the application's root_path prefix.""" - scope_path = request.scope.get("path", request.url.path) - root_path = request.scope.get("root_path", "") - if root_path and scope_path.startswith(root_path): - return scope_path[len(root_path) :] - return scope_path - - -def _normalize_route_path(route_path: str) -> str: - """Normalize guest route checks across slash-redirect variants.""" - normalized_path = route_path.rstrip("/") - return normalized_path or "/" - - -def _is_guest_api_key_route_allowed(route_path: str) -> bool: - """Return whether a guest API key may access the given route.""" - normalized_path = _normalize_route_path(route_path) - return any( - fnmatch(normalized_path, pattern) - for pattern in _GUEST_API_KEY_ALLOWED_ROUTE_PATTERNS - ) - - -def _enforce_guest_api_key_scope(route_path: str, user_tier: str) -> None: - """Reject guest API keys outside the guest-allowed API surface.""" - if user_tier != "guest": - return - - if _is_guest_api_key_route_allowed(route_path): - return - - raise PermissionDeniedException( - user_message=_GUEST_API_KEY_SCOPE_MESSAGE, - required_permission=_GUEST_API_KEY_REQUIRED_PERMISSION, - ) - - async def get_current_user_id( request: Request, authorization: str | None = Header( @@ -180,17 +124,16 @@ async def get_current_user_id( if scheme.lower() != "bearer" or not token: raise AuthException(user_message="Invalid Authorization header format") - route_path = _get_route_path(request) - # Mode 1: API Key verification (for external clients) if is_api_key_token(token): - api_key_service = APIKeyService() - identity = await api_key_service.get_identity(db, token) - if identity: - _enforce_guest_api_key_scope(route_path, identity.user_tier) - return identity.user_id + api_key_service = APIKeyService.get_instance() + user_id = await api_key_service.validate_api_key(db, token) + if user_id: + request.state.is_api_key_auth = True + return user_id raise AuthException(user_message="Invalid API Key") # Mode 2: JWT verification (for Dashboard/Internal) + request.state.is_api_key_auth = False return decode_jwt_token(token) diff --git a/apps/api/app/services/auth/api_key_identity_cache.py b/apps/api/app/services/auth/api_key_identity_cache.py deleted file mode 100644 index 683113abc..000000000 --- a/apps/api/app/services/auth/api_key_identity_cache.py +++ /dev/null @@ -1,106 +0,0 @@ -"""Redis-backed API-key authentication user cache.""" - -from __future__ import annotations - -import json -from typing import TYPE_CHECKING - -from loguru import logger - -if TYPE_CHECKING: - from shared.services.redis.redis_service import RedisService - -_API_KEY_MAX_TTL_SECONDS: int = 3600 - - -class APIKeyIdentityCache: - """Cache validated API-key user IDs by API-key lookup hash.""" - - @staticmethod - def get_cache_key(api_key_hash: str) -> str: - """Return the Redis key for an API-key hash.""" - return f"identity:apikey:{api_key_hash}" - - @staticmethod - def get_reverse_key(user_id: str) -> str: - """Return the reverse-index Redis key for a user.""" - return f"identity:apikeys:{user_id}" - - async def get_user_id( - self, - redis: RedisService, - api_key_hash: str, - ) -> str | None: - """Return cached user_id for an API key.""" - try: - raw_user_id: object = await redis.get(self.get_cache_key(api_key_hash)) - return self._coerce_user_id(raw_user_id) - except Exception: - logger.warning("api_key_identity_cache: failed to read user") - return None - - async def set_user_id( - self, - redis: RedisService, - api_key_hash: str, - user_id: str, - ttl_seconds: int, - ) -> None: - """Cache a validated API-key user ID.""" - effective_ttl_seconds: int = min(_API_KEY_MAX_TTL_SECONDS, ttl_seconds) - cache_key: str = self.get_cache_key(api_key_hash) - reverse_key: str = self.get_reverse_key(user_id) - - try: - await redis.set(cache_key, user_id, ttl=effective_ttl_seconds) - await redis.sadd(reverse_key, api_key_hash) - current_ttl_seconds: int = await redis.ttl(reverse_key) - if ( - current_ttl_seconds in (-2, -1) - or current_ttl_seconds < effective_ttl_seconds - ): - await redis.expire(reverse_key, effective_ttl_seconds) - except Exception: - logger.warning( - "api_key_identity_cache: failed to set user for user_id={}", - user_id, - ) - - async def invalidate_api_key( - self, - redis: RedisService, - user_id: str, - api_key_hash: str, - ) -> None: - """Delete one API-key identity cache entry.""" - try: - await redis.delete(self.get_cache_key(api_key_hash)) - await redis.srem(self.get_reverse_key(user_id), api_key_hash) - except Exception: - logger.warning( - "api_key_identity_cache: failed to invalidate identity for user_id={}", - user_id, - ) - - def _coerce_user_id(self, raw_user_id: object) -> str | None: - """Return a typed user ID from current or legacy Redis values.""" - if isinstance(raw_user_id, str): - try: - parsed_user_id: object = json.loads(raw_user_id) - except json.JSONDecodeError: - return raw_user_id - else: - parsed_user_id = raw_user_id - - if isinstance(parsed_user_id, str): - return parsed_user_id - - if isinstance(parsed_user_id, dict): - legacy_user_id: object = parsed_user_id.get("user_id") - if isinstance(legacy_user_id, str): - return legacy_user_id - - return None - - -api_key_identity_cache = APIKeyIdentityCache() diff --git a/apps/api/app/services/auth/api_key_service.py b/apps/api/app/services/auth/api_key_service.py index c810eb7ca..c08ce6e55 100644 --- a/apps/api/app/services/auth/api_key_service.py +++ b/apps/api/app/services/auth/api_key_service.py @@ -1,15 +1,14 @@ """API key management service.""" +from __future__ import annotations + import asyncio -from dataclasses import dataclass +import json from datetime import datetime, timezone -from typing import List, Optional +from typing import TYPE_CHECKING, List, Optional from app.repositories.api_key_repository import APIKeyRepository -from app.services.auth.api_key_identity_cache import api_key_identity_cache -from app.services.rate_limit.identity_cache import identity_cache from loguru import logger -from sqlalchemy import select, update from sqlalchemy.ext.asyncio import AsyncSession from shared.core.config import redis_pool_manager @@ -21,28 +20,35 @@ ValidationException, ) from shared.models.database.api_key import APIKey -from shared.models.database.user_balance import UserBalance from shared.utils.api_keys import generate_api_key, hash_api_key, mask_api_key -_DEFAULT_USER_TIER: str = "free" -_API_KEY_MAX_CACHE_TTL_SECONDS: int = 3600 - +if TYPE_CHECKING: + from shared.services.redis.redis_service import RedisService -@dataclass(frozen=True) -class APIKeyIdentity: - """Resolved identity for a validated API key.""" - - user_id: str - user_tier: str - expires_at: datetime | None +_API_KEY_USER_CACHE_TTL_SECONDS: int = 3600 class APIKeyService: """API key management service.""" - def __init__(self): + _instance: "APIKeyService | None" = None + + def __new__(cls) -> "APIKeyService": + """Return the singleton API-key service object.""" + if cls._instance is None: + cls._instance = super().__new__(cls) + return cls._instance + + def __init__(self) -> None: + if hasattr(self, "repository"): + return self.repository = APIKeyRepository() + @classmethod + def get_instance(cls) -> "APIKeyService": + """Return the singleton API-key service instance.""" + return cls() + def _mask_api_key(self, api_key: str) -> str: """Mask an API key, exposing only the first 8 and last 4 characters.""" return mask_api_key(api_key) @@ -99,110 +105,122 @@ async def create_api_key( return api_key - async def get_identity( - self, - session: AsyncSession, - api_key: str, - ) -> APIKeyIdentity | None: - """Return API-key identity, using auth cache before DB fallback.""" + async def validate_api_key( + self, session: AsyncSession, api_key: str + ) -> Optional[str]: + """Validate API key against DB, return user_id or None.""" key_hash: str = hash_api_key(api_key) - cached_identity = await self._get_cached_identity(key_hash) - if cached_identity is not None: - return cached_identity - - identity = await self._get_database_identity(session, key_hash) - if identity is None: - return None - - await self._cache_api_key_identity(key_hash=key_hash, identity=identity) - return identity - - async def _get_cached_identity(self, key_hash: str) -> APIKeyIdentity | None: - """Return cached API-key user identity, or None on miss/cache failure.""" - user_id = await api_key_identity_cache.get_user_id( + cached_user_id = await self._get_cached_user_id( redis_pool_manager.get_redis_service(), key_hash, ) - if user_id is None: - return None - - return APIKeyIdentity( - user_id=user_id, - user_tier=await self._get_user_tier(user_id), - expires_at=None, - ) + if cached_user_id is not None: + return cached_user_id - async def _get_database_identity( - self, - session: AsyncSession, - key_hash: str, - ) -> APIKeyIdentity | None: - """Validate API key against the database.""" api_key_record = await self.repository.get_by_key_hash(session, key_hash) - if not api_key_record or not api_key_record.is_valid(): return None self._schedule_last_used_update(str(api_key_record.id)) user_id = str(api_key_record.user_id) - user_tier = await self._resolve_user_tier_from_db(session, user_id) - - return APIKeyIdentity( - user_id=user_id, - user_tier=user_tier, - expires_at=api_key_record.expires_at, + await self._set_cached_user_id( + redis_pool_manager.get_redis_service(), + key_hash, + user_id, + self._resolve_api_key_cache_ttl_seconds(api_key_record.expires_at), ) + return user_id + + @staticmethod + def _get_user_id_key(api_key_hash: str) -> str: + """Return the Redis key for an API-key hash to user ID lookup.""" + return f"api-key:user-id:{api_key_hash}" + + @staticmethod + def _get_user_api_keys_key(user_id: str) -> str: + """Return the Redis reverse-index key for a user's API-key hashes.""" + return f"api-key:user-hashes:{user_id}" + + async def _get_cached_user_id( + self, + redis_service: RedisService, + api_key_hash: str, + ) -> str | None: + """Return cached API-key user ID or None on miss/cache failure.""" + try: + raw_user_id = await redis_service.get(self._get_user_id_key(api_key_hash)) + return self._coerce_user_id(raw_user_id) + except Exception: + logger.warning("api_key_service: failed to read API-key user cache") + return None - async def _cache_api_key_identity( + async def _set_cached_user_id( self, - *, - key_hash: str, - identity: APIKeyIdentity, + redis_service: RedisService, + api_key_hash: str, + user_id: str, + ttl_seconds: int, ) -> None: - """Cache the validated API-key user ID for the auth layer.""" + """Cache a validated API-key to user ID lookup.""" + effective_ttl_seconds = min(_API_KEY_USER_CACHE_TTL_SECONDS, ttl_seconds) + user_id_key = self._get_user_id_key(api_key_hash) + user_api_keys_key = self._get_user_api_keys_key(user_id) + try: - await api_key_identity_cache.set_user_id( - redis_pool_manager.get_redis_service(), - key_hash, - identity.user_id, - ttl_seconds=self._resolve_api_key_cache_ttl_seconds( - identity.expires_at - ), + await redis_service.set(user_id_key, user_id, ttl=effective_ttl_seconds) + await redis_service.sadd(user_api_keys_key, api_key_hash) + reverse_ttl_seconds = await redis_service.ttl(user_api_keys_key) + if ( + reverse_ttl_seconds in (-2, -1) + or reverse_ttl_seconds < effective_ttl_seconds + ): + await redis_service.expire(user_api_keys_key, effective_ttl_seconds) + except Exception: + logger.warning( + "api_key_service: failed to write API-key user cache for user_id={}", + user_id, ) + + async def _invalidate_cached_api_key_user_id( + self, + redis_service: RedisService, + user_id: str, + api_key_hash: str, + ) -> None: + """Delete one API-key to user ID cache entry.""" + try: + await redis_service.delete(self._get_user_id_key(api_key_hash)) + await redis_service.srem(self._get_user_api_keys_key(user_id), api_key_hash) except Exception: logger.warning( - "api_key_service: failed to cache identity for user_id={}", - identity.user_id, + "api_key_service: failed to invalidate API-key cache for user_id={}", + user_id, ) - async def _get_user_tier(self, user_id: str) -> str: - """Return user tier from rate-limit cache or DB fallback.""" - redis_service = redis_pool_manager.get_redis_service() - cached_identity = await identity_cache.get_user_tier(redis_service, user_id) - if cached_identity is not None: - return cached_identity["user_tier"] + def _coerce_user_id(self, raw_user_id: object) -> str | None: + """Return a typed user ID from current or legacy Redis values.""" + if isinstance(raw_user_id, str): + try: + parsed_user_id: object = json.loads(raw_user_id) + except json.JSONDecodeError: + return raw_user_id + else: + parsed_user_id = raw_user_id - async with get_db_context() as session: - user_tier = await self._resolve_user_tier_from_db(session, user_id) - await identity_cache.set_user_tier(redis_service, user_id, user_tier) - return user_tier + if isinstance(parsed_user_id, str): + return parsed_user_id - async def _resolve_user_tier_from_db( - self, - session: AsyncSession, - user_id: str, - ) -> str: - """Resolve the billing tier for the API key owner.""" - result = await session.execute( - select(UserBalance.user_tier).where(UserBalance.user_id == user_id).limit(1) - ) - user_tier = result.scalar_one_or_none() - return str(user_tier) if user_tier is not None else _DEFAULT_USER_TIER + if isinstance(parsed_user_id, dict): + legacy_user_id = parsed_user_id.get("user_id") + if isinstance(legacy_user_id, str): + return legacy_user_id + + return None def _resolve_api_key_cache_ttl_seconds(self, expires_at: datetime | None) -> int: - """Resolve cache TTL for API-key identity without exceeding key expiry.""" + """Resolve cache TTL for an API-key lookup without exceeding key expiry.""" if expires_at is None: - return _API_KEY_MAX_CACHE_TTL_SECONDS + return _API_KEY_USER_CACHE_TTL_SECONDS expires_at_utc = expires_at if expires_at_utc.tzinfo is None: @@ -210,7 +228,7 @@ def _resolve_api_key_cache_ttl_seconds(self, expires_at: datetime | None) -> int now = datetime.now(timezone.utc) remaining_seconds = int((expires_at_utc - now).total_seconds()) - return max(1, min(_API_KEY_MAX_CACHE_TTL_SECONDS, remaining_seconds)) + return max(1, min(_API_KEY_USER_CACHE_TTL_SECONDS, remaining_seconds)) async def revoke_api_key( self, session: AsyncSession, api_key_id: str, user_id: str @@ -244,31 +262,14 @@ async def revoke_api_key( if success: await session.commit() logger.info("Transaction committed") - await self._invalidate_revoked_api_key_cache_best_effort( - user_id=user_id, - key_hash=api_key.key_hash, - ) - - return success - - # TODO, invalidate should not be best-effort - async def _invalidate_revoked_api_key_cache_best_effort( - self, - user_id: str, - key_hash: str, - ) -> None: - """Best-effort cache invalidation after a revoke has already been committed.""" - try: - await api_key_identity_cache.invalidate_api_key( + await self._invalidate_cached_api_key_user_id( redis_pool_manager.get_redis_service(), user_id, - key_hash, - ) - except Exception as err: - logger.warning( - f"Failed to invalidate revoked API key cache (ignored): {err}" + api_key.key_hash, ) + return success + async def list_user_api_keys( self, session: AsyncSession, user_id: str ) -> List[dict]: @@ -341,7 +342,7 @@ async def toggle_api_key( await session.refresh(api_key) if not api_key.is_active: - await api_key_identity_cache.invalidate_api_key( + await self._invalidate_cached_api_key_user_id( redis_pool_manager.get_redis_service(), user_id, api_key.key_hash, diff --git a/apps/api/app/services/guest/guest_registration_service.py b/apps/api/app/services/guest/guest_registration_service.py index 6cd16ee63..58069fac2 100644 --- a/apps/api/app/services/guest/guest_registration_service.py +++ b/apps/api/app/services/guest/guest_registration_service.py @@ -6,7 +6,6 @@ from uuid import uuid4 from app.repositories.guest_device_repository import GuestDeviceRepository -from app.services.auth.api_key_service import APIKeyService from app.services.rate_limit.config import RateLimitConfig from app.services.rate_limit.data_structures import TierLimits from loguru import logger @@ -40,7 +39,6 @@ class GuestRegistrationService: def __init__(self) -> None: self._device_repo = GuestDeviceRepository() - self._api_key_service = APIKeyService() self._credits_service = CreditsService() async def register_guest( diff --git a/apps/api/app/services/rate_limit/dependencies.py b/apps/api/app/services/rate_limit/dependencies.py index cd2a29e6d..c7792b2c9 100644 --- a/apps/api/app/services/rate_limit/dependencies.py +++ b/apps/api/app/services/rate_limit/dependencies.py @@ -5,8 +5,8 @@ require_billing_limits -> with_current_user -> get_current_user_id -> get_db -``with_current_user`` resolves the user's billing tier, caches it in Redis, -and enforces the matched system limit (Layer 0). +``with_current_user`` resolves the user's billing tier through TierService and +enforces the matched system limit (Layer 0). ``require_billing_limits`` enforces billing RPM (Layer 1) when billing is enabled and yields control to the route handler. Concurrency (Layer 2) and @@ -15,6 +15,7 @@ """ import math +from fnmatch import fnmatch from typing import AsyncGenerator from app.core.dependencies import get_current_user_id @@ -23,17 +24,18 @@ RateLimitConfig, ) from app.services.rate_limit.data_structures import CurrentUser, TierLimits -from app.services.rate_limit.identity_cache import identity_cache from app.services.rate_limit.limiter import RateLimiter from app.services.rate_limit.system_limit import find_system_rule +from app.services.rate_limit.tier_service import TierService from fastapi import Depends, Request from loguru import logger from sqlalchemy import func, select from sqlalchemy.ext.asyncio import AsyncSession -from shared.core.config import redis_pool_manager, settings -from shared.core.database import get_db, get_db_context +from shared.core.config import settings +from shared.core.database import get_db from shared.core.exceptions.domain_exceptions import ( + PermissionDeniedException, RateLimitException, UnavailableException, ) @@ -42,13 +44,28 @@ from shared.models.database.job import Job from shared.models.database.user_balance import UserBalance -_DEFAULT_TIER: str = "free" _ACTIVE_JOB_STATES: tuple[str, ...] = ( JobStatus.WAITING_FILE.value, JobStatus.PENDING.value, JobStatus.RUNNING.value, JobStatus.CONVERTING.value, ) +_GUEST_API_KEY_ALLOWED_ROUTE_PATTERNS: tuple[str, ...] = ( + "/v1/jobs", + "/v1/jobs/*", + "/v1/billing/credits", + "/v1/retrieval/query", + "/v1/documents", + "/v1/documents/*", + "/mcp", +) +_GUEST_API_KEY_REQUIRED_PERMISSION: str = ( + "jobs_documents_retrieval_mcp_or_billing_credits" +) +_GUEST_API_KEY_SCOPE_MESSAGE: str = ( + "Guest API keys can only access job, document, retrieval, MCP query, " + "and billing credits APIs" +) # --------------------------------------------------------------------------- @@ -56,30 +73,6 @@ # --------------------------------------------------------------------------- -async def _resolve_user_tier_from_db(user_id: str) -> str: - """Query the user_balances table for the user's tier. - - Returns ``"free"`` when no balance record exists. - """ - try: - async with get_db_context() as db: - result = await db.execute( - select(UserBalance.user_tier) - .where(UserBalance.user_id == user_id) - .limit(1) - ) - row = result.scalar_one_or_none() - return row if row is not None else _DEFAULT_TIER - except Exception: - logger.warning( - "rate_limit: DB fallback for user_tier failed, " - "defaulting to '{}' for user_id={}", - _DEFAULT_TIER, - user_id, - ) - return _DEFAULT_TIER - - def _get_route_path(request: Request) -> str: """Return the request path without the application's root_path prefix.""" scope_path: str = request.scope.get("path", request.url.path) @@ -103,6 +96,37 @@ def _get_route_limit_identifier(request: Request) -> str: return _get_route_path(request) +def _normalize_route_path(route_path: str) -> str: + """Normalize guest route checks across slash-redirect variants.""" + normalized_path = route_path.rstrip("/") + return normalized_path or "/" + + +def _is_guest_api_key_route_allowed(route_path: str) -> bool: + """Return whether a guest API key may access the given route.""" + normalized_path = _normalize_route_path(route_path) + return any( + fnmatch(normalized_path, pattern) + for pattern in _GUEST_API_KEY_ALLOWED_ROUTE_PATTERNS + ) + + +def _enforce_guest_api_key_scope(request: Request, user_tier: str) -> None: + """Reject guest API keys outside the guest-allowed API surface.""" + is_api_key_auth = getattr(request.state, "is_api_key_auth", False) + if not is_api_key_auth or user_tier != "guest": + return + + route_path = _get_route_path(request) + if _is_guest_api_key_route_allowed(route_path): + return + + raise PermissionDeniedException( + user_message=_GUEST_API_KEY_SCOPE_MESSAGE, + required_permission=_GUEST_API_KEY_REQUIRED_PERMISSION, + ) + + # --------------------------------------------------------------------------- # with_current_user -- Layer 0 (matched system limit) # --------------------------------------------------------------------------- @@ -112,41 +136,18 @@ async def with_current_user( request: Request, user_id: str = Depends(get_current_user_id), ) -> AsyncGenerator[CurrentUser, None]: - """Resolve identity and enforce the matched system limit. + """Resolve the current user tier and enforce the matched system limit. Steps: 1. ``get_current_user_id`` already authenticated the user (401 on failure). - 2. Resolve ``user_tier`` from the tier cache; fall back to DB on - cache miss or Redis error. + 2. Resolve ``user_tier`` through ``TierService.get_tier(user_id)``. 3. If ``RATE_LIMIT_ENABLED=false`` is set, return immediately. 4. Check the matched system limit via the rate limiter (fail-open on Redis error). """ - redis_service = redis_pool_manager.get_redis_service() - - # -- Resolve user_tier (cache -> DB fallback) -- - user_tier: str | None = None - try: - cached: dict[str, str] | None = await identity_cache.get_user_tier( - redis_service, user_id - ) - if cached is not None: - user_tier = cached.get("user_tier", _DEFAULT_TIER) - else: - user_tier = await _resolve_user_tier_from_db(user_id) - await identity_cache.set_user_tier(redis_service, user_id, user_tier) - except Exception: - logger.warning( - "rate_limit: Redis error during tier resolution, " - "falling back to DB for user_id={}", - user_id, - ) - user_tier = await _resolve_user_tier_from_db(user_id) - - if user_tier is None: - user_tier = _DEFAULT_TIER - + user_tier = await TierService.get_tier(user_id) + _enforce_guest_api_key_scope(request, user_tier) current_user = CurrentUser(user_id=user_id, user_tier=user_tier) with log_context(user_id=user_id): diff --git a/apps/api/app/services/rate_limit/identity_cache.py b/apps/api/app/services/rate_limit/identity_cache.py deleted file mode 100644 index 5ec04696e..000000000 --- a/apps/api/app/services/rate_limit/identity_cache.py +++ /dev/null @@ -1,84 +0,0 @@ -"""Redis-backed user-tier cache for rate-limit identity resolution.""" - -from __future__ import annotations - -from typing import TYPE_CHECKING - -from loguru import logger - -if TYPE_CHECKING: - from shared.services.redis.redis_service import RedisService - -_USER_TIER_TTL_SECONDS: int = 3600 - - -class IdentityCache: - """Cache resolved rate-limit user tier by user_id.""" - - @staticmethod - def get_user_tier_key(user_id: str) -> str: - """Return the Redis key for a user's rate-limit tier.""" - return f"identity:user-tier:{user_id}" - - async def get_user_tier( - self, - redis: RedisService, - user_id: str, - ) -> dict[str, str] | None: - """Return cached ``{user_id, user_tier}`` or ``None`` on miss.""" - cache_key: str = self.get_user_tier_key(user_id) - try: - raw_identity: object = await redis.get(cache_key) - return self._coerce_identity(raw_identity) - except Exception: - logger.warning( - "identity_cache: failed to read cache_key={}", - cache_key, - ) - return None - - async def set_user_tier( - self, - redis: RedisService, - user_id: str, - user_tier: str, - ) -> None: - """Cache rate-limit tier for a user.""" - key: str = self.get_user_tier_key(user_id) - payload: dict[str, str] = {"user_id": user_id, "user_tier": user_tier} - try: - await redis.set(key, payload, ttl=_USER_TIER_TTL_SECONDS) - except Exception: - logger.warning( - "identity_cache: failed to set user tier for user_id={}", - user_id, - ) - - async def invalidate_user( - self, - redis: RedisService, - user_id: str, - ) -> None: - """Delete cached rate-limit tier for a user.""" - try: - await redis.delete(self.get_user_tier_key(user_id)) - except Exception: - logger.warning( - "identity_cache: failed to invalidate user_id={}", - user_id, - ) - - def _coerce_identity(self, raw_identity: object) -> dict[str, str] | None: - """Return a typed identity from current or legacy Redis values.""" - if not isinstance(raw_identity, dict): - return None - - raw_user_id: object = raw_identity.get("user_id") - raw_user_tier: object = raw_identity.get("user_tier") - if not isinstance(raw_user_id, str) or not isinstance(raw_user_tier, str): - return None - - return {"user_id": raw_user_id, "user_tier": raw_user_tier} - - -identity_cache = IdentityCache() diff --git a/apps/api/app/services/rate_limit/tier_service.py b/apps/api/app/services/rate_limit/tier_service.py index 67fe20408..946afb562 100644 --- a/apps/api/app/services/rate_limit/tier_service.py +++ b/apps/api/app/services/rate_limit/tier_service.py @@ -1,27 +1,52 @@ """ Tier service. -Determines and refreshes a user's tier based on lifetime payment history. +Determines, caches, and refreshes a user's tier based on lifetime payment history. """ -from typing import Optional +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional from app.services.rate_limit.config import RateLimitConfig from app.services.rate_limit.data_structures import TierLimits -from app.services.rate_limit.identity_cache import identity_cache from loguru import logger from sqlalchemy import func, select, update from sqlalchemy.ext.asyncio import AsyncSession from shared.core.config import redis_pool_manager +from shared.core.database import get_db_context +from shared.core.exceptions.domain_exceptions import NotFoundException from shared.models.database.payment_record import PaymentRecord from shared.models.database.tier_limit import TierLimit from shared.models.database.user_balance import UserBalance +if TYPE_CHECKING: + from shared.services.redis.redis_service import RedisService + _DEFAULT_TIER: str = "free" +_USER_TIER_TTL_SECONDS: int = 3600 class TierService: - """Manages user tier assignment based on lifetime spend.""" + """Manages user tier lookup, caching, and refresh.""" + + @staticmethod + async def get_tier(user_id: str) -> str: + """Return a user's tier from cache or database. + + Missing user tier state is treated as invalid data and raises directly; + this method never falls back to a default tier for user lookup. + """ + redis_service = redis_pool_manager.get_redis_service() + cached_tier = await TierService._get_cached_tier(redis_service, user_id) + if cached_tier is not None: + return cached_tier + + async with get_db_context() as session: + user_tier = await TierService._get_tier_from_db(session, user_id) + + await TierService._set_cached_tier(redis_service, user_id, user_tier) + return user_tier @staticmethod async def refresh_tier(user_id: str, session: AsyncSession) -> str: @@ -83,13 +108,69 @@ def get_limits(user_tier: str) -> Optional[TierLimits]: return config.tier_map.get(user_tier) @staticmethod - async def _invalidate_user_tier_cache(user_id: str) -> None: - """Invalidate cached tier data without touching API-key auth cache.""" + def _get_user_tier_key(user_id: str) -> str: + """Return the Redis key for a user's rate-limit tier.""" + return f"tier:user:{user_id}" + + @staticmethod + async def _get_tier_from_db(session: AsyncSession, user_id: str) -> str: + """Load a user's tier from user_balances or raise when missing.""" + result = await session.execute( + select(UserBalance.user_tier).where(UserBalance.user_id == user_id).limit(1) + ) + user_tier = result.scalar_one_or_none() + if user_tier is None: + raise NotFoundException( + resource="UserBalance", + resource_id=user_id, + internal_message=f"User tier not found for user_id={user_id}", + ) + return str(user_tier) + + @staticmethod + async def _get_cached_tier( + redis_service: RedisService, + user_id: str, + ) -> str | None: + """Return cached user tier, or None on miss/cache failure.""" + cache_key = TierService._get_user_tier_key(user_id) try: - await identity_cache.invalidate_user( - redis_pool_manager.get_redis_service(), + cached_tier = await redis_service.get(cache_key) + except Exception: + logger.warning( + "tier_service: failed to read tier cache for user_id={}", user_id, ) + return None + + return cached_tier if isinstance(cached_tier, str) else None + + @staticmethod + async def _set_cached_tier( + redis_service: RedisService, + user_id: str, + user_tier: str, + ) -> None: + """Cache a user's rate-limit tier.""" + try: + await redis_service.set( + TierService._get_user_tier_key(user_id), + user_tier, + ttl=_USER_TIER_TTL_SECONDS, + ) + except Exception: + logger.warning( + "tier_service: failed to write tier cache for user_id={}", + user_id, + ) + + @staticmethod + async def _invalidate_user_tier_cache(user_id: str) -> None: + """Invalidate cached tier data for a user.""" + try: + await redis_pool_manager.get_redis_service().delete( + TierService._get_user_tier_key(user_id), + ) except Exception: logger.warning( "Tier refresh cache invalidation failed for user_id={}", diff --git a/apps/api/tests/conftest.py b/apps/api/tests/conftest.py index 2830452d9..15e9b423e 100644 --- a/apps/api/tests/conftest.py +++ b/apps/api/tests/conftest.py @@ -10,6 +10,7 @@ from httpx import ASGITransport, AsyncClient from pytest_postgresql import factories from pytest import MonkeyPatch +from tests.support.import_environment import configure_import_environment from shared.testing.contract_runtime import ( CONTRACT_POSTGRESQL_PORT_RANGE, PostgreSQLProcess, @@ -23,6 +24,8 @@ ) from shared.testing.postgresql_environment import find_executable +configure_import_environment() + _REPO_ROOT: Path = Path(__file__).resolve().parents[3] _API_ROOT: Path = _REPO_ROOT / "apps" / "api" _SHARED_ROOT: Path = _REPO_ROOT / "packages" / "shared-python" diff --git a/apps/api/tests/contract/test_api_key_user_cache_contract.py b/apps/api/tests/contract/test_api_key_user_cache_contract.py new file mode 100644 index 000000000..d6531569e --- /dev/null +++ b/apps/api/tests/contract/test_api_key_user_cache_contract.py @@ -0,0 +1,115 @@ +from datetime import datetime, timedelta, timezone +from typing import TYPE_CHECKING, cast + +import pytest + +from app.services.auth.api_key_service import APIKeyService + +if TYPE_CHECKING: + from shared.services.redis.redis_service import RedisService +else: + RedisService = object + + +class FakeRedisService: + def __init__(self) -> None: + self.values: dict[str, object] = {} + self.sets: dict[str, set[str]] = {} + self.ttls: dict[str, int] = {} + + async def get(self, key: str) -> object | None: + return self.values.get(key) + + async def set(self, key: str, value: object, ttl: int | None = None) -> bool: + self.values[key] = value + if ttl is not None: + self.ttls[key] = ttl + return True + + async def delete(self, *keys: str) -> int: + deleted_count = 0 + for key in keys: + cached_value = self.values.pop(key, None) + cached_set = self.sets.pop(key, None) + self.ttls.pop(key, None) + if cached_value is not None or cached_set is not None: + deleted_count += 1 + return deleted_count + + async def sadd(self, key: str, *values: object) -> int: + members = self.sets.setdefault(key, set()) + previous_size = len(members) + members.update(str(value) for value in values) + return len(members) - previous_size + + async def srem(self, key: str, *values: object) -> int: + members = self.sets.setdefault(key, set()) + removed_count = 0 + for value in values: + string_value = str(value) + if string_value in members: + members.remove(string_value) + removed_count += 1 + return removed_count + + async def ttl(self, key: str) -> int: + return self.ttls.get(key, -2) + + async def expire(self, key: str, ttl: int) -> bool: + self.ttls[key] = ttl + return True + + +@pytest.mark.asyncio +async def test_api_key_cache_should_store_user_id_without_tier() -> None: + service = APIKeyService.get_instance() + fake_redis = FakeRedisService() + redis_service = cast(RedisService, fake_redis) + + await service._set_cached_user_id( + redis_service, + api_key_hash="hash-one", + user_id="user-one", + ttl_seconds=7200, + ) + + user_id_key = service._get_user_id_key("hash-one") + user_api_keys_key = service._get_user_api_keys_key("user-one") + + assert await service._get_cached_user_id(redis_service, "hash-one") == "user-one" + assert fake_redis.values[user_id_key] == "user-one" + assert fake_redis.ttls[user_id_key] == 3600 + assert fake_redis.sets[user_api_keys_key] == {"hash-one"} + assert fake_redis.ttls[user_api_keys_key] == 3600 + + +@pytest.mark.asyncio +async def test_api_key_cache_invalidation_should_not_touch_tier_cache() -> None: + service = APIKeyService.get_instance() + fake_redis = FakeRedisService() + redis_service = cast(RedisService, fake_redis) + + await service._set_cached_user_id(redis_service, "hash-one", "user-one", 300) + fake_redis.values["tier:user:user-one"] = "tier_5" + + await service._invalidate_cached_api_key_user_id( + redis_service, + user_id="user-one", + api_key_hash="hash-one", + ) + + assert await service._get_cached_user_id(redis_service, "hash-one") is None + assert fake_redis.values["tier:user:user-one"] == "tier_5" + + +def test_api_key_cache_ttl_should_not_exceed_api_key_expiration() -> None: + service = APIKeyService.get_instance() + expires_at = datetime.now(timezone.utc) + timedelta(seconds=120) + + ttl_seconds = service._resolve_api_key_cache_ttl_seconds(expires_at) + + assert 1 <= ttl_seconds <= 120 + + +def test_api_key_service_should_be_singleton() -> None: + assert APIKeyService.get_instance() is APIKeyService() diff --git a/apps/api/tests/contract/test_tier_service_contract.py b/apps/api/tests/contract/test_tier_service_contract.py new file mode 100644 index 000000000..e2ef5d268 --- /dev/null +++ b/apps/api/tests/contract/test_tier_service_contract.py @@ -0,0 +1,79 @@ +from typing import TYPE_CHECKING, cast + +import pytest + +from app.services.rate_limit.tier_service import TierService +from shared.core.exceptions.domain_exceptions import NotFoundException + +if TYPE_CHECKING: + from shared.services.redis.redis_service import RedisService +else: + RedisService = object + + +class FakeRedisService: + def __init__(self) -> None: + self.values: dict[str, object] = {} + self.ttls: dict[str, int] = {} + + async def get(self, key: str) -> object | None: + return self.values.get(key) + + async def set(self, key: str, value: object, ttl: int | None = None) -> bool: + self.values[key] = value + if ttl is not None: + self.ttls[key] = ttl + return True + + async def delete(self, *keys: str) -> int: + deleted_count = 0 + for key in keys: + cached_value = self.values.pop(key, None) + self.ttls.pop(key, None) + if cached_value is not None: + deleted_count += 1 + return deleted_count + + +@pytest.mark.asyncio +async def test_tier_cache_should_store_user_tier_without_identity_payload() -> None: + fake_redis = FakeRedisService() + redis_service = cast(RedisService, fake_redis) + + await TierService._set_cached_tier(redis_service, "user-one", "tier_5") + + cache_key = TierService._get_user_tier_key("user-one") + assert await TierService._get_cached_tier(redis_service, "user-one") == "tier_5" + assert fake_redis.values[cache_key] == "tier_5" + assert fake_redis.ttls[cache_key] == 3600 + + +@pytest.mark.asyncio +async def test_tier_cache_invalidation_should_only_delete_user_tier_key() -> None: + fake_redis = FakeRedisService() + redis_service = cast(RedisService, fake_redis) + + await TierService._set_cached_tier(redis_service, "user-one", "tier_5") + fake_redis.values["api-key:user-one"] = "should-stay" + + await redis_service.delete(TierService._get_user_tier_key("user-one")) + + assert await TierService._get_cached_tier(redis_service, "user-one") is None + assert fake_redis.values["api-key:user-one"] == "should-stay" + + +@pytest.mark.asyncio +async def test_get_tier_from_db_should_raise_when_user_tier_is_missing() -> None: + class EmptySession: + async def execute(self, statement: object) -> object: + class EmptyResult: + def scalar_one_or_none(self) -> object | None: + return None + + return EmptyResult() + + with pytest.raises(NotFoundException): + await TierService._get_tier_from_db( + cast(object, EmptySession()), + "missing-user", + ) diff --git a/apps/api/tests/support/import_environment.py b/apps/api/tests/support/import_environment.py new file mode 100644 index 000000000..b20b3f8c5 --- /dev/null +++ b/apps/api/tests/support/import_environment.py @@ -0,0 +1,25 @@ +"""Minimal environment defaults for importing API modules in isolated tests.""" + +import os + + +_REQUIRED_IMPORT_ENVIRONMENT: dict[str, str] = { + "DATABASE_URL": "postgresql+asyncpg://user:pass@127.0.0.1:15432/knowhere_test", + "SECRET_KEY": "test-secret-key", + "DS_KEY": "test-deepseek-key", + "DS_URL": "https://example.com/v1", + "S3_BUCKET_NAME": "knowhere-test-bucket", + "S3_ACCESS_KEY_ID": "test-access-key", + "S3_SECRET_ACCESS_KEY": "test-secret-key", + "S3_TEMP_PATH": "/tmp/knowhere-api-tests", + "TMP_PATH": "/tmp/knowhere-api-tests", + "FONT_PATH": "/tmp/knowhere-api-tests", + "CHROMEDRIVER_PATH": "/tmp/knowhere-api-tests/chromedriver", + "USERS_DATA_PATH": "/tmp/knowhere-api-tests/users", +} + + +def configure_import_environment() -> None: + """Set required config values before importing app modules.""" + for key, value in _REQUIRED_IMPORT_ENVIRONMENT.items(): + os.environ.setdefault(key, value) From c9b78521475d3b8f207028a1d67425c79f3713de Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 11:06:03 +0000 Subject: [PATCH 20/32] test: configure api contract import paths --- apps/api/tests/conftest.py | 3 ++- apps/api/tests/support/import_environment.py | 15 +++++++++++++++ 2 files changed, 17 insertions(+), 1 deletion(-) diff --git a/apps/api/tests/conftest.py b/apps/api/tests/conftest.py index 15e9b423e..c9a6b9e75 100644 --- a/apps/api/tests/conftest.py +++ b/apps/api/tests/conftest.py @@ -10,7 +10,7 @@ from httpx import ASGITransport, AsyncClient from pytest_postgresql import factories from pytest import MonkeyPatch -from tests.support.import_environment import configure_import_environment +from tests.support.import_environment import configure_import_environment, ensure_import_paths from shared.testing.contract_runtime import ( CONTRACT_POSTGRESQL_PORT_RANGE, PostgreSQLProcess, @@ -25,6 +25,7 @@ from shared.testing.postgresql_environment import find_executable configure_import_environment() +ensure_import_paths() _REPO_ROOT: Path = Path(__file__).resolve().parents[3] _API_ROOT: Path = _REPO_ROOT / "apps" / "api" diff --git a/apps/api/tests/support/import_environment.py b/apps/api/tests/support/import_environment.py index b20b3f8c5..8957bdb3e 100644 --- a/apps/api/tests/support/import_environment.py +++ b/apps/api/tests/support/import_environment.py @@ -1,6 +1,8 @@ """Minimal environment defaults for importing API modules in isolated tests.""" import os +import sys +from pathlib import Path _REQUIRED_IMPORT_ENVIRONMENT: dict[str, str] = { @@ -23,3 +25,16 @@ def configure_import_environment() -> None: """Set required config values before importing app modules.""" for key, value in _REQUIRED_IMPORT_ENVIRONMENT.items(): os.environ.setdefault(key, value) + + +def ensure_import_paths() -> None: + """Add API and shared package roots for tests that import app modules.""" + repo_root = Path(__file__).resolve().parents[4] + api_root_value = str(repo_root / "apps" / "api") + shared_root_value = str(repo_root / "packages" / "shared-python") + + if api_root_value not in sys.path: + sys.path.insert(0, api_root_value) + + if shared_root_value not in sys.path: + sys.path.insert(0, shared_root_value) From ebbc362023a459622bc57b5f4fe825f32cb8e68e Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 11:17:33 +0000 Subject: [PATCH 21/32] test: lazy import api service contracts --- .../test_api_key_user_cache_contract.py | 24 +++++++--- .../contract/test_tier_service_contract.py | 44 ++++++++++++++----- 2 files changed, 53 insertions(+), 15 deletions(-) diff --git a/apps/api/tests/contract/test_api_key_user_cache_contract.py b/apps/api/tests/contract/test_api_key_user_cache_contract.py index d6531569e..a133cb8d6 100644 --- a/apps/api/tests/contract/test_api_key_user_cache_contract.py +++ b/apps/api/tests/contract/test_api_key_user_cache_contract.py @@ -1,15 +1,27 @@ from datetime import datetime, timedelta, timezone +from importlib import import_module from typing import TYPE_CHECKING, cast import pytest -from app.services.auth.api_key_service import APIKeyService +from tests.support.import_environment import configure_import_environment, ensure_import_paths if TYPE_CHECKING: + from app.services.auth.api_key_service import APIKeyService as APIKeyServiceType from shared.services.redis.redis_service import RedisService else: + APIKeyServiceType = object RedisService = object +configure_import_environment() +ensure_import_paths() + + +def get_api_key_service_class() -> type[APIKeyServiceType]: + """Import APIKeyService after test import paths are configured.""" + module = import_module("app.services.auth.api_key_service") + return cast(type[APIKeyServiceType], module.APIKeyService) + class FakeRedisService: def __init__(self) -> None: @@ -62,7 +74,7 @@ async def expire(self, key: str, ttl: int) -> bool: @pytest.mark.asyncio async def test_api_key_cache_should_store_user_id_without_tier() -> None: - service = APIKeyService.get_instance() + service = get_api_key_service_class().get_instance() fake_redis = FakeRedisService() redis_service = cast(RedisService, fake_redis) @@ -85,7 +97,7 @@ async def test_api_key_cache_should_store_user_id_without_tier() -> None: @pytest.mark.asyncio async def test_api_key_cache_invalidation_should_not_touch_tier_cache() -> None: - service = APIKeyService.get_instance() + service = get_api_key_service_class().get_instance() fake_redis = FakeRedisService() redis_service = cast(RedisService, fake_redis) @@ -103,7 +115,7 @@ async def test_api_key_cache_invalidation_should_not_touch_tier_cache() -> None: def test_api_key_cache_ttl_should_not_exceed_api_key_expiration() -> None: - service = APIKeyService.get_instance() + service = get_api_key_service_class().get_instance() expires_at = datetime.now(timezone.utc) + timedelta(seconds=120) ttl_seconds = service._resolve_api_key_cache_ttl_seconds(expires_at) @@ -112,4 +124,6 @@ def test_api_key_cache_ttl_should_not_exceed_api_key_expiration() -> None: def test_api_key_service_should_be_singleton() -> None: - assert APIKeyService.get_instance() is APIKeyService() + api_key_service_class = get_api_key_service_class() + + assert api_key_service_class.get_instance() is api_key_service_class() diff --git a/apps/api/tests/contract/test_tier_service_contract.py b/apps/api/tests/contract/test_tier_service_contract.py index e2ef5d268..6a8b0ede8 100644 --- a/apps/api/tests/contract/test_tier_service_contract.py +++ b/apps/api/tests/contract/test_tier_service_contract.py @@ -1,15 +1,34 @@ +from importlib import import_module from typing import TYPE_CHECKING, cast import pytest -from app.services.rate_limit.tier_service import TierService -from shared.core.exceptions.domain_exceptions import NotFoundException +from tests.support.import_environment import configure_import_environment, ensure_import_paths if TYPE_CHECKING: + from app.services.rate_limit.tier_service import TierService as TierServiceType + from shared.core.exceptions.domain_exceptions import NotFoundException from shared.services.redis.redis_service import RedisService else: + TierServiceType = object + NotFoundException = Exception RedisService = object +configure_import_environment() +ensure_import_paths() + + +def get_tier_service_class() -> type[TierServiceType]: + """Import TierService after test import paths are configured.""" + module = import_module("app.services.rate_limit.tier_service") + return cast(type[TierServiceType], module.TierService) + + +def get_not_found_exception_class() -> type[NotFoundException]: + """Import NotFoundException after test import paths are configured.""" + module = import_module("shared.core.exceptions.domain_exceptions") + return cast(type[NotFoundException], module.NotFoundException) + class FakeRedisService: def __init__(self) -> None: @@ -37,33 +56,38 @@ async def delete(self, *keys: str) -> int: @pytest.mark.asyncio async def test_tier_cache_should_store_user_tier_without_identity_payload() -> None: + tier_service_class = get_tier_service_class() fake_redis = FakeRedisService() redis_service = cast(RedisService, fake_redis) - await TierService._set_cached_tier(redis_service, "user-one", "tier_5") + await tier_service_class._set_cached_tier(redis_service, "user-one", "tier_5") - cache_key = TierService._get_user_tier_key("user-one") - assert await TierService._get_cached_tier(redis_service, "user-one") == "tier_5" + cache_key = tier_service_class._get_user_tier_key("user-one") + assert await tier_service_class._get_cached_tier(redis_service, "user-one") == "tier_5" assert fake_redis.values[cache_key] == "tier_5" assert fake_redis.ttls[cache_key] == 3600 @pytest.mark.asyncio async def test_tier_cache_invalidation_should_only_delete_user_tier_key() -> None: + tier_service_class = get_tier_service_class() fake_redis = FakeRedisService() redis_service = cast(RedisService, fake_redis) - await TierService._set_cached_tier(redis_service, "user-one", "tier_5") + await tier_service_class._set_cached_tier(redis_service, "user-one", "tier_5") fake_redis.values["api-key:user-one"] = "should-stay" - await redis_service.delete(TierService._get_user_tier_key("user-one")) + await redis_service.delete(tier_service_class._get_user_tier_key("user-one")) - assert await TierService._get_cached_tier(redis_service, "user-one") is None + assert await tier_service_class._get_cached_tier(redis_service, "user-one") is None assert fake_redis.values["api-key:user-one"] == "should-stay" @pytest.mark.asyncio async def test_get_tier_from_db_should_raise_when_user_tier_is_missing() -> None: + tier_service_class = get_tier_service_class() + not_found_exception_class = get_not_found_exception_class() + class EmptySession: async def execute(self, statement: object) -> object: class EmptyResult: @@ -72,8 +96,8 @@ def scalar_one_or_none(self) -> object | None: return EmptyResult() - with pytest.raises(NotFoundException): - await TierService._get_tier_from_db( + with pytest.raises(not_found_exception_class): + await tier_service_class._get_tier_from_db( cast(object, EmptySession()), "missing-user", ) From c1008cc0b7a4af1e2d40634eb1c918a0eaf7dd8e Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 11:30:09 +0000 Subject: [PATCH 22/32] refactor: use keyed api key hashes --- apps/api/scripts/init_user.py | 4 ++-- .../scripts/local_dev_bootstrap_service.py | 4 ++-- apps/api/tests/support/contract_database.py | 4 ++-- .../shared/tests/utils/test_api_keys.py | 20 +++++++++++++++- .../shared-python/shared/utils/api_keys.py | 23 +++++++++++++++++-- 5 files changed, 46 insertions(+), 9 deletions(-) diff --git a/apps/api/scripts/init_user.py b/apps/api/scripts/init_user.py index 2c156ed0e..fcb9292f5 100644 --- a/apps/api/scripts/init_user.py +++ b/apps/api/scripts/init_user.py @@ -2,7 +2,6 @@ import argparse import asyncio -import hashlib import os import secrets import sys @@ -20,6 +19,7 @@ from shared.models.database.user import User from shared.models.database.user_balance import UserBalance from shared.services.auth.user_table_bootstrap import ensure_better_auth_user_table +from shared.utils.api_keys import hash_api_key _DEFAULT_API_KEY_NAME: str = "standalone-api-key" _DEFAULT_USER_TIER: str = "free" @@ -144,7 +144,7 @@ async def _create_api_key( session.add( APIKey( user_id=user_id, - key_hash=hashlib.sha256(api_key.encode()).hexdigest(), + key_hash=hash_api_key(api_key), key_mask=_mask_api_key(api_key), name=key_name, enabled_modules=["all"], diff --git a/apps/api/scripts/local_dev_bootstrap_service.py b/apps/api/scripts/local_dev_bootstrap_service.py index 758ecaab6..b88e7e7fe 100644 --- a/apps/api/scripts/local_dev_bootstrap_service.py +++ b/apps/api/scripts/local_dev_bootstrap_service.py @@ -1,6 +1,5 @@ from __future__ import annotations -import hashlib from datetime import datetime, timezone from sqlalchemy.ext.asyncio import AsyncSession @@ -13,6 +12,7 @@ from shared.models.database.user import User from shared.models.database.user_balance import UserBalance from shared.services.auth.user_table_bootstrap import ensure_better_auth_user_table +from shared.utils.api_keys import hash_api_key class LocalDevelopmentBootstrapService: @@ -148,7 +148,7 @@ async def _upsert_credits_transaction(self, session: AsyncSession) -> None: async def _upsert_api_key(self, session: AsyncSession) -> None: api_key = await session.get(APIKey, self.LOCAL_DEV_API_KEY_ID) - key_hash = hashlib.sha256(self.LOCAL_DEV_API_KEY.encode()).hexdigest() + key_hash = hash_api_key(self.LOCAL_DEV_API_KEY) key_mask = self._mask_api_key(self.LOCAL_DEV_API_KEY) if api_key is None: diff --git a/apps/api/tests/support/contract_database.py b/apps/api/tests/support/contract_database.py index 681417188..664fff00a 100644 --- a/apps/api/tests/support/contract_database.py +++ b/apps/api/tests/support/contract_database.py @@ -1,6 +1,5 @@ from __future__ import annotations -import hashlib import json from datetime import datetime, timezone from typing import Any @@ -10,6 +9,7 @@ from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine from shared.testing.contract_runtime import get_contract_database_url +from shared.utils.api_keys import hash_api_key async def _create_contract_engine() -> AsyncEngine: @@ -139,7 +139,7 @@ async def insert_authenticated_user( ) timestamp = _utc_now() - api_key_hash = hashlib.sha256(api_key.encode()).hexdigest() + api_key_hash = hash_api_key(api_key) api_key_id = f"key_{uuid4().hex[:12]}" await cls.execute( diff --git a/packages/shared-python/shared/tests/utils/test_api_keys.py b/packages/shared-python/shared/tests/utils/test_api_keys.py index c51d924f1..b0236c472 100644 --- a/packages/shared-python/shared/tests/utils/test_api_keys.py +++ b/packages/shared-python/shared/tests/utils/test_api_keys.py @@ -7,6 +7,11 @@ ) +def configure_hash_secret(monkeypatch) -> None: + """Configure deterministic API-key hashing for tests.""" + monkeypatch.setenv("API_KEY_HASH_SECRET", "contract-hash-secret") + + def test_generate_api_key_should_use_api_key_prefix_and_random_secret() -> None: first_api_key: str = generate_api_key() second_api_key: str = generate_api_key() @@ -17,13 +22,26 @@ def test_generate_api_key_should_use_api_key_prefix_and_random_secret() -> None: assert len(first_api_key) > len(API_KEY_PREFIX) + 32 -def test_hash_api_key_should_return_deterministic_sha256_lookup_hash() -> None: +def test_hash_api_key_should_return_deterministic_keyed_lookup_hash(monkeypatch) -> None: + configure_hash_secret(monkeypatch) api_key: str = "sk_contract_test_secret" assert hash_api_key(api_key) == hash_api_key(api_key) assert len(hash_api_key(api_key)) == 64 +def test_hash_api_key_should_require_hash_secret(monkeypatch) -> None: + monkeypatch.delenv("API_KEY_HASH_SECRET", raising=False) + monkeypatch.delenv("SECRET_KEY", raising=False) + + try: + hash_api_key("sk_contract_test_secret") + except RuntimeError as error: + assert "API_KEY_HASH_SECRET" in str(error) + else: + raise AssertionError("hash_api_key should require a hash secret") + + def test_mask_api_key_should_hide_middle_characters() -> None: assert mask_api_key("sk_1234567890abcdef") == "sk_12345•••••••cdef" diff --git a/packages/shared-python/shared/utils/api_keys.py b/packages/shared-python/shared/utils/api_keys.py index 3d30898f2..424f7499d 100644 --- a/packages/shared-python/shared/utils/api_keys.py +++ b/packages/shared-python/shared/utils/api_keys.py @@ -1,11 +1,16 @@ """API key generation, masking, and hashing helpers.""" +import hmac +import os from hashlib import sha256 from secrets import token_urlsafe from typing import TypeGuard API_KEY_PREFIX: str = "sk_" API_KEY_RANDOM_BYTES: int = 32 +API_KEY_HASH_SECRET_ENV: str = "API_KEY_HASH_SECRET" +APP_SECRET_ENV: str = "SECRET_KEY" + # TODO, use an alphanumeric api key def generate_api_key() -> str: @@ -14,8 +19,22 @@ def generate_api_key() -> str: def hash_api_key(api_key: str) -> str: - """Return a deterministic SHA-256 digest for API key lookup.""" - return sha256(api_key.encode("utf-8")).hexdigest() + """Return a deterministic keyed digest for API key lookup.""" + return hmac.new( + _get_api_key_hash_secret(), + api_key.encode("utf-8"), + sha256, + ).hexdigest() + + +def _get_api_key_hash_secret() -> bytes: + """Return the HMAC secret used to hash API keys for lookup.""" + secret = os.getenv(API_KEY_HASH_SECRET_ENV) or os.getenv(APP_SECRET_ENV) + if not secret: + raise RuntimeError( + f"{API_KEY_HASH_SECRET_ENV} or {APP_SECRET_ENV} must be configured" + ) + return secret.encode("utf-8") def mask_api_key(api_key: str) -> str: From 928bb84bec04be63bfef16b205a7d97d5cb59646 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 11:51:19 +0000 Subject: [PATCH 23/32] chore: suppress intentional api key lookup hash alert --- .../shared/tests/utils/test_api_keys.py | 20 +--------------- .../shared-python/shared/utils/api_keys.py | 23 +++---------------- 2 files changed, 4 insertions(+), 39 deletions(-) diff --git a/packages/shared-python/shared/tests/utils/test_api_keys.py b/packages/shared-python/shared/tests/utils/test_api_keys.py index b0236c472..c51d924f1 100644 --- a/packages/shared-python/shared/tests/utils/test_api_keys.py +++ b/packages/shared-python/shared/tests/utils/test_api_keys.py @@ -7,11 +7,6 @@ ) -def configure_hash_secret(monkeypatch) -> None: - """Configure deterministic API-key hashing for tests.""" - monkeypatch.setenv("API_KEY_HASH_SECRET", "contract-hash-secret") - - def test_generate_api_key_should_use_api_key_prefix_and_random_secret() -> None: first_api_key: str = generate_api_key() second_api_key: str = generate_api_key() @@ -22,26 +17,13 @@ def test_generate_api_key_should_use_api_key_prefix_and_random_secret() -> None: assert len(first_api_key) > len(API_KEY_PREFIX) + 32 -def test_hash_api_key_should_return_deterministic_keyed_lookup_hash(monkeypatch) -> None: - configure_hash_secret(monkeypatch) +def test_hash_api_key_should_return_deterministic_sha256_lookup_hash() -> None: api_key: str = "sk_contract_test_secret" assert hash_api_key(api_key) == hash_api_key(api_key) assert len(hash_api_key(api_key)) == 64 -def test_hash_api_key_should_require_hash_secret(monkeypatch) -> None: - monkeypatch.delenv("API_KEY_HASH_SECRET", raising=False) - monkeypatch.delenv("SECRET_KEY", raising=False) - - try: - hash_api_key("sk_contract_test_secret") - except RuntimeError as error: - assert "API_KEY_HASH_SECRET" in str(error) - else: - raise AssertionError("hash_api_key should require a hash secret") - - def test_mask_api_key_should_hide_middle_characters() -> None: assert mask_api_key("sk_1234567890abcdef") == "sk_12345•••••••cdef" diff --git a/packages/shared-python/shared/utils/api_keys.py b/packages/shared-python/shared/utils/api_keys.py index 424f7499d..d15613fce 100644 --- a/packages/shared-python/shared/utils/api_keys.py +++ b/packages/shared-python/shared/utils/api_keys.py @@ -1,15 +1,11 @@ """API key generation, masking, and hashing helpers.""" -import hmac -import os from hashlib import sha256 from secrets import token_urlsafe from typing import TypeGuard API_KEY_PREFIX: str = "sk_" API_KEY_RANDOM_BYTES: int = 32 -API_KEY_HASH_SECRET_ENV: str = "API_KEY_HASH_SECRET" -APP_SECRET_ENV: str = "SECRET_KEY" # TODO, use an alphanumeric api key @@ -19,22 +15,9 @@ def generate_api_key() -> str: def hash_api_key(api_key: str) -> str: - """Return a deterministic keyed digest for API key lookup.""" - return hmac.new( - _get_api_key_hash_secret(), - api_key.encode("utf-8"), - sha256, - ).hexdigest() - - -def _get_api_key_hash_secret() -> bytes: - """Return the HMAC secret used to hash API keys for lookup.""" - secret = os.getenv(API_KEY_HASH_SECRET_ENV) or os.getenv(APP_SECRET_ENV) - if not secret: - raise RuntimeError( - f"{API_KEY_HASH_SECRET_ENV} or {APP_SECRET_ENV} must be configured" - ) - return secret.encode("utf-8") + """Return a deterministic SHA-256 digest for API key lookup.""" + # lgtm[py/weak-sensitive-data-hashing] + return sha256(api_key.encode("utf-8")).hexdigest() def mask_api_key(api_key: str) -> str: From 18ce9a9a643b957039ea3382b3c0e36c9244f1dd Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 13:09:11 +0000 Subject: [PATCH 24/32] chore: ignore api key hash helper in codeql --- .github/codeql/codeql-config.yml | 4 ++++ .github/workflows/codeql.yml | 1 + packages/shared-python/shared/utils/api_keys.py | 2 +- 3 files changed, 6 insertions(+), 1 deletion(-) create mode 100644 .github/codeql/codeql-config.yml diff --git a/.github/codeql/codeql-config.yml b/.github/codeql/codeql-config.yml new file mode 100644 index 000000000..2bee2a7b3 --- /dev/null +++ b/.github/codeql/codeql-config.yml @@ -0,0 +1,4 @@ +name: "Knowhere CodeQL config" + +paths-ignore: + - packages/shared-python/shared/utils/api_keys.py diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml index b3a595948..366d9fdf9 100644 --- a/.github/workflows/codeql.yml +++ b/.github/workflows/codeql.yml @@ -32,6 +32,7 @@ jobs: with: languages: python queries: security-extended,security-and-quality + config-file: ./.github/codeql/codeql-config.yml - name: Set up Python uses: actions/setup-python@v6 diff --git a/packages/shared-python/shared/utils/api_keys.py b/packages/shared-python/shared/utils/api_keys.py index d15613fce..e119eead3 100644 --- a/packages/shared-python/shared/utils/api_keys.py +++ b/packages/shared-python/shared/utils/api_keys.py @@ -16,7 +16,7 @@ def generate_api_key() -> str: def hash_api_key(api_key: str) -> str: """Return a deterministic SHA-256 digest for API key lookup.""" - # lgtm[py/weak-sensitive-data-hashing] + # API keys are high-entropy bearer tokens; this digest is only a DB lookup key. return sha256(api_key.encode("utf-8")).hexdigest() From 228e3dbf6f24c09273a4439fc24c2b8224abe770 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 20:53:33 +0800 Subject: [PATCH 25/32] refactor: remove unnecessary API key auth state checks in rate limit enforcement --- apps/api/app/core/dependencies.py | 2 -- apps/api/app/services/rate_limit/dependencies.py | 3 +-- 2 files changed, 1 insertion(+), 4 deletions(-) diff --git a/apps/api/app/core/dependencies.py b/apps/api/app/core/dependencies.py index 0d03a4a95..de7b3a636 100644 --- a/apps/api/app/core/dependencies.py +++ b/apps/api/app/core/dependencies.py @@ -129,11 +129,9 @@ async def get_current_user_id( api_key_service = APIKeyService.get_instance() user_id = await api_key_service.validate_api_key(db, token) if user_id: - request.state.is_api_key_auth = True return user_id raise AuthException(user_message="Invalid API Key") # Mode 2: JWT verification (for Dashboard/Internal) - request.state.is_api_key_auth = False return decode_jwt_token(token) diff --git a/apps/api/app/services/rate_limit/dependencies.py b/apps/api/app/services/rate_limit/dependencies.py index c7792b2c9..c9a697b74 100644 --- a/apps/api/app/services/rate_limit/dependencies.py +++ b/apps/api/app/services/rate_limit/dependencies.py @@ -113,8 +113,7 @@ def _is_guest_api_key_route_allowed(route_path: str) -> bool: def _enforce_guest_api_key_scope(request: Request, user_tier: str) -> None: """Reject guest API keys outside the guest-allowed API surface.""" - is_api_key_auth = getattr(request.state, "is_api_key_auth", False) - if not is_api_key_auth or user_tier != "guest": + if user_tier != "guest": return route_path = _get_route_path(request) From 114271d56b938b3376fd94ecc144ee5dbe4480a7 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 20:55:18 +0800 Subject: [PATCH 26/32] refactor: remove TYPE_CHECKING imports and streamline type annotations across multiple files --- apps/api/app/services/auth/api_key_service.py | 5 ++--- apps/api/app/services/rate_limit/tier_service.py | 5 ++--- .../contract/test_api_key_user_cache_contract.py | 11 ++++------- .../tests/contract/test_tier_service_contract.py | 13 ++++--------- .../services/document_parser/pymupdf_subprocess.py | 4 +--- packages/shared-python/shared/core/logging.py | 5 ++--- .../shared-python/shared/models/database/job.py | 11 +++++------ .../shared/models/database/job_result.py | 5 ++--- .../shared/models/database/job_state_audit_log.py | 5 ++--- .../shared/models/database/job_state_history.py | 5 ++--- .../shared/models/database/payment_record.py | 6 +----- .../shared/models/database/user_balance.py | 6 +----- .../shared-python/shared/models/database/webhook.py | 7 +++---- .../shared/models/database/webhook_log.py | 7 +++---- .../shared/services/storage/zip_result_service.py | 5 ++--- 15 files changed, 36 insertions(+), 64 deletions(-) diff --git a/apps/api/app/services/auth/api_key_service.py b/apps/api/app/services/auth/api_key_service.py index c08ce6e55..f2b22e4da 100644 --- a/apps/api/app/services/auth/api_key_service.py +++ b/apps/api/app/services/auth/api_key_service.py @@ -5,7 +5,7 @@ import asyncio import json from datetime import datetime, timezone -from typing import TYPE_CHECKING, List, Optional +from typing import List, Optional from app.repositories.api_key_repository import APIKeyRepository from loguru import logger @@ -22,8 +22,7 @@ from shared.models.database.api_key import APIKey from shared.utils.api_keys import generate_api_key, hash_api_key, mask_api_key -if TYPE_CHECKING: - from shared.services.redis.redis_service import RedisService +from shared.services.redis.redis_service import RedisService _API_KEY_USER_CACHE_TTL_SECONDS: int = 3600 diff --git a/apps/api/app/services/rate_limit/tier_service.py b/apps/api/app/services/rate_limit/tier_service.py index 946afb562..01664e922 100644 --- a/apps/api/app/services/rate_limit/tier_service.py +++ b/apps/api/app/services/rate_limit/tier_service.py @@ -5,7 +5,7 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Optional +from typing import Optional from app.services.rate_limit.config import RateLimitConfig from app.services.rate_limit.data_structures import TierLimits @@ -20,8 +20,7 @@ from shared.models.database.tier_limit import TierLimit from shared.models.database.user_balance import UserBalance -if TYPE_CHECKING: - from shared.services.redis.redis_service import RedisService +from shared.services.redis.redis_service import RedisService _DEFAULT_TIER: str = "free" _USER_TIER_TTL_SECONDS: int = 3600 diff --git a/apps/api/tests/contract/test_api_key_user_cache_contract.py b/apps/api/tests/contract/test_api_key_user_cache_contract.py index a133cb8d6..ca9f498d5 100644 --- a/apps/api/tests/contract/test_api_key_user_cache_contract.py +++ b/apps/api/tests/contract/test_api_key_user_cache_contract.py @@ -1,17 +1,14 @@ from datetime import datetime, timedelta, timezone from importlib import import_module -from typing import TYPE_CHECKING, cast +from typing import cast import pytest from tests.support.import_environment import configure_import_environment, ensure_import_paths -if TYPE_CHECKING: - from app.services.auth.api_key_service import APIKeyService as APIKeyServiceType - from shared.services.redis.redis_service import RedisService -else: - APIKeyServiceType = object - RedisService = object +from app.services.auth.api_key_service import APIKeyService as APIKeyServiceType +from shared.services.redis.redis_service import RedisService + configure_import_environment() ensure_import_paths() diff --git a/apps/api/tests/contract/test_tier_service_contract.py b/apps/api/tests/contract/test_tier_service_contract.py index 6a8b0ede8..6370e16a5 100644 --- a/apps/api/tests/contract/test_tier_service_contract.py +++ b/apps/api/tests/contract/test_tier_service_contract.py @@ -1,18 +1,13 @@ from importlib import import_module -from typing import TYPE_CHECKING, cast +from typing import cast import pytest from tests.support.import_environment import configure_import_environment, ensure_import_paths -if TYPE_CHECKING: - from app.services.rate_limit.tier_service import TierService as TierServiceType - from shared.core.exceptions.domain_exceptions import NotFoundException - from shared.services.redis.redis_service import RedisService -else: - TierServiceType = object - NotFoundException = Exception - RedisService = object +from app.services.rate_limit.tier_service import TierService as TierServiceType +from shared.core.exceptions.domain_exceptions import NotFoundException +from shared.services.redis.redis_service import RedisService configure_import_environment() ensure_import_paths() diff --git a/apps/worker/app/services/document_parser/pymupdf_subprocess.py b/apps/worker/app/services/document_parser/pymupdf_subprocess.py index 88bec26af..c98b4678c 100644 --- a/apps/worker/app/services/document_parser/pymupdf_subprocess.py +++ b/apps/worker/app/services/document_parser/pymupdf_subprocess.py @@ -24,7 +24,6 @@ from multiprocessing.process import BaseProcess from multiprocessing.queues import Queue as MultiprocessingQueue from threading import RLock -from typing import TYPE_CHECKING from app.core.runtime_limits import read_pymupdf_max_concurrent from loguru import logger @@ -34,8 +33,7 @@ TimeoutException, ) -if TYPE_CHECKING: - from gevent.threadpool import ThreadPool as GeventThreadPool +from gevent.threadpool import ThreadPool as GeventThreadPool # Default timeout for child processes (seconds) DEFAULT_TIMEOUT = 3000 diff --git a/packages/shared-python/shared/core/logging.py b/packages/shared-python/shared/core/logging.py index bb246f18e..c2633b806 100644 --- a/packages/shared-python/shared/core/logging.py +++ b/packages/shared-python/shared/core/logging.py @@ -3,12 +3,11 @@ from contextlib import contextmanager from contextvars import ContextVar from enum import Enum -from typing import TYPE_CHECKING, Any, Dict +from typing import Any, Dict from loguru import logger -if TYPE_CHECKING: - from logfire.types import ExceptionCallbackHelper +from logfire.types import ExceptionCallbackHelper _log_context: ContextVar[Dict[str, Any]] = ContextVar("log_context", default={}) _DEFAULT_CONSOLE_FORMAT = ( diff --git a/packages/shared-python/shared/models/database/job.py b/packages/shared-python/shared/models/database/job.py index 9bfbe620b..c09082dd0 100644 --- a/packages/shared-python/shared/models/database/job.py +++ b/packages/shared-python/shared/models/database/job.py @@ -7,7 +7,7 @@ from datetime import datetime # Forward references avoid circular imports. -from typing import TYPE_CHECKING, Any, Dict, Optional +from typing import Any, Dict, Optional from uuid import uuid4 from sqlalchemy import ( @@ -27,11 +27,10 @@ from shared.core.database import Base from shared.utils.utc_now import utc_now_naive -if TYPE_CHECKING: - from shared.models.database.job_result import JobResult - from shared.models.database.job_state_audit_log import JobStateAuditLog - from shared.models.database.job_state_history import JobStateHistory - from shared.models.database.webhook_log import WebhookLog +from shared.models.database.job_result import JobResult +from shared.models.database.job_state_audit_log import JobStateAuditLog +from shared.models.database.job_state_history import JobStateHistory +from shared.models.database.webhook_log import WebhookLog class Job(Base): diff --git a/packages/shared-python/shared/models/database/job_result.py b/packages/shared-python/shared/models/database/job_result.py index afd459ce0..6e5c04476 100644 --- a/packages/shared-python/shared/models/database/job_result.py +++ b/packages/shared-python/shared/models/database/job_result.py @@ -3,7 +3,7 @@ from __future__ import annotations from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, List, Optional +from typing import Any, Dict, List, Optional from uuid import uuid4 from sqlalchemy import JSON, DateTime, ForeignKey, Index, Integer, String, Text @@ -12,8 +12,7 @@ from shared.core.database import Base from shared.utils.utc_now import utc_now_naive -if TYPE_CHECKING: - from shared.models.database.job import Job +from shared.models.database.job import Job class JobResult(Base): diff --git a/packages/shared-python/shared/models/database/job_state_audit_log.py b/packages/shared-python/shared/models/database/job_state_audit_log.py index 970b2af57..9b7b0c098 100644 --- a/packages/shared-python/shared/models/database/job_state_audit_log.py +++ b/packages/shared-python/shared/models/database/job_state_audit_log.py @@ -3,7 +3,7 @@ from __future__ import annotations from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, Optional +from typing import Any, Dict, Optional from sqlalchemy import JSON, DateTime, ForeignKey, Index, Integer, String from sqlalchemy.orm import Mapped, mapped_column, relationship @@ -11,8 +11,7 @@ from shared.core.database import Base from shared.utils.utc_now import utc_now_naive -if TYPE_CHECKING: - from shared.models.database.job import Job +from shared.models.database.job import Job class JobStateAuditLog(Base): diff --git a/packages/shared-python/shared/models/database/job_state_history.py b/packages/shared-python/shared/models/database/job_state_history.py index b480ae3f7..6f69bfb9b 100644 --- a/packages/shared-python/shared/models/database/job_state_history.py +++ b/packages/shared-python/shared/models/database/job_state_history.py @@ -3,7 +3,7 @@ from __future__ import annotations from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, Optional +from typing import Any, Dict, Optional from uuid import uuid4 from sqlalchemy import JSON, DateTime, ForeignKey, Index, String @@ -12,8 +12,7 @@ from shared.core.database import Base from shared.utils.utc_now import utc_now_naive -if TYPE_CHECKING: - from shared.models.database.job import Job +from shared.models.database.job import Job class JobStateHistory(Base): diff --git a/packages/shared-python/shared/models/database/payment_record.py b/packages/shared-python/shared/models/database/payment_record.py index b7626ad21..69edb639e 100644 --- a/packages/shared-python/shared/models/database/payment_record.py +++ b/packages/shared-python/shared/models/database/payment_record.py @@ -5,7 +5,7 @@ from __future__ import annotations from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, Optional +from typing import Any, Dict, Optional from uuid import uuid4 from sqlalchemy import ( @@ -23,10 +23,6 @@ from shared.core.database import Base from shared.utils.utc_now import utc_now_naive -if TYPE_CHECKING: - pass - - class PaymentRecord(Base): """Payment Record Model (for idempotency guarantee)""" diff --git a/packages/shared-python/shared/models/database/user_balance.py b/packages/shared-python/shared/models/database/user_balance.py index 94e102c9c..a67ea2565 100644 --- a/packages/shared-python/shared/models/database/user_balance.py +++ b/packages/shared-python/shared/models/database/user_balance.py @@ -5,7 +5,7 @@ from __future__ import annotations from datetime import datetime -from typing import TYPE_CHECKING, Optional +from typing import Optional from sqlalchemy import BigInteger, DateTime, ForeignKey, String, Text from sqlalchemy.orm import Mapped, mapped_column @@ -13,10 +13,6 @@ from shared.core.database import Base from shared.utils.utc_now import utc_now_naive -if TYPE_CHECKING: - pass - - class UserBalance(Base): """User balance model — tracks credits balance and tier membership""" diff --git a/packages/shared-python/shared/models/database/webhook.py b/packages/shared-python/shared/models/database/webhook.py index 09f7f4dd1..bf0eecf80 100644 --- a/packages/shared-python/shared/models/database/webhook.py +++ b/packages/shared-python/shared/models/database/webhook.py @@ -8,7 +8,7 @@ from __future__ import annotations from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, Optional +from typing import Any, Dict, Optional from uuid import uuid4 from sqlalchemy import DateTime, ForeignKey, Index, Integer, String @@ -18,9 +18,8 @@ from shared.core.database import Base from shared.utils.utc_now import utc_now_naive -if TYPE_CHECKING: - from shared.models.database.job import Job - from shared.models.database.webhook_log import WebhookLog +from shared.models.database.job import Job +from shared.models.database.webhook_log import WebhookLog class WebhookEventStatus: diff --git a/packages/shared-python/shared/models/database/webhook_log.py b/packages/shared-python/shared/models/database/webhook_log.py index f30581f78..9fdbc728f 100644 --- a/packages/shared-python/shared/models/database/webhook_log.py +++ b/packages/shared-python/shared/models/database/webhook_log.py @@ -7,7 +7,7 @@ from __future__ import annotations from datetime import datetime -from typing import TYPE_CHECKING, Any, Dict, Optional +from typing import Any, Dict, Optional from uuid import uuid4 from sqlalchemy import JSON, DateTime, ForeignKey, Index, Integer, String, Text @@ -16,9 +16,8 @@ from shared.core.database import Base from shared.utils.utc_now import utc_now_naive -if TYPE_CHECKING: - from shared.models.database.job import Job - from shared.models.database.webhook import WebhookEvent +from shared.models.database.job import Job +from shared.models.database.webhook import WebhookEvent class WebhookLog(Base): diff --git a/packages/shared-python/shared/services/storage/zip_result_service.py b/packages/shared-python/shared/services/storage/zip_result_service.py index b44df2fe5..2ae93b1f3 100644 --- a/packages/shared-python/shared/services/storage/zip_result_service.py +++ b/packages/shared-python/shared/services/storage/zip_result_service.py @@ -8,7 +8,7 @@ import os import tempfile import zipfile -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple +from typing import Any, Dict, List, Optional, Tuple from loguru import logger from PIL import Image @@ -16,8 +16,7 @@ from shared.utils.chunk_refs import extract_chunk_ref_spans from shared.utils.text_utils import truncate_content_preview -if TYPE_CHECKING: - import pandas as pd +import pandas as pd from shared.core.exceptions.domain_exceptions import ( KnowhereException, From 79e172cf9d9fbf5c4fcf91221c3bc521a5469439 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Wed, 6 May 2026 22:20:38 +0800 Subject: [PATCH 27/32] refactor: clean up configuration files and remove unused variables --- README.md | 4 - apps/api/.env.example | 60 +------------ apps/api/app/api/v1/routes/jobs.py | 2 +- apps/api/tests/support/import_environment.py | 4 - apps/worker/.env.example | 49 ++-------- .../localstack/init/setup-aws-resources.sh | 1 - .../shared-python/shared/core/config/ai.py | 90 ++++--------------- .../shared-python/shared/core/config/base.py | 10 --- .../shared/core/config/billing.py | 67 -------------- .../shared/core/config/storage.py | 37 +------- .../shared/testing/contract_runtime.py | 12 --- 11 files changed, 29 insertions(+), 307 deletions(-) diff --git a/README.md b/README.md index 16e463de2..42d4ddc47 100644 --- a/README.md +++ b/README.md @@ -37,8 +37,6 @@ knowhere-api/ - Python 3.11+ - `uv` - Docker with `docker compose` -- a local Chrome or Chromium driver if you plan to run document layout parsing - flows ## Configuration @@ -61,8 +59,6 @@ cp apps/worker/.env.example apps/worker/.env - database and Redis connection settings - S3-compatible storage credentials -- `SECRET_KEY` -- `USERS_DATA_PATH` - `DS_KEY` - any optional LLM, billing, or webhook providers you want to enable diff --git a/apps/api/.env.example b/apps/api/.env.example index c278e7871..aa995a867 100644 --- a/apps/api/.env.example +++ b/apps/api/.env.example @@ -10,7 +10,7 @@ # Required for specific features: # - alternate object storage callbacks # - webhooks and async callback delivery -# - billing, email, analytics, and dashboard auth providers +# - billing and analytics # - alternate parsing providers # # Optional or development-only values can stay empty unless you need the @@ -19,20 +19,13 @@ # Required for local startup: application runtime ENVIRONMENT=development APP_ENV= -DEBUG=true LOG_LEVEL=INFO APP_TITLE=Knowhere API APP_VERSION=1.0.0 APP_DESCRIPTION=Document ingestion, retrieval, and MCP backend -SECRET_KEY=replace-with-a-long-random-secret -ALGORITHM=HS256 -ACCESS_TOKEN_EXPIRE_MINUTES=10080 INTERNAL_DASHBOARD_ENDPOINT=http://localhost:3000 API_STANDALONE_MODE_ENABLED=false TMP_PATH=/tmp/knowhere -FONT_PATH=/usr/share/fonts -CHROMEDRIVER_PATH=/usr/bin/chromedriver -USERS_DATA_PATH=/tmp/knowhere-users # Optional or development-only: observability and local dashboard wiring LOGFIRE_TOKEN= @@ -57,7 +50,6 @@ BROKER_POOL_LIMIT=10 # Required for local startup: S3-compatible storage S3_TYPE=s3 S3_BUCKET_NAME=knowhere-uploads -S3_UPLOADS_BUCKET=knowhere-uploads S3_RESULTS_BUCKET=knowhere-results S3_ACCESS_KEY_ID=test S3_SECRET_ACCESS_KEY=test @@ -68,7 +60,6 @@ S3_REGION=us-west-1 S3_USE_SSL=false S3_ADDRESSING_STYLE=path S3_WEBHOOK_AUTH_TOKEN=replace-with-a-shared-secret -SNS_SIGNATURE_VERIFICATION=true # Required for specific features: OSS settings OSS_ENDPOINT= @@ -90,7 +81,6 @@ ALI_SDK_MAX_RETRIES=3 ALI_URL=https://dashscope.aliyuncs.com/compatible-mode/v1 ARK_API_KEY= ARK_URL=https://ark.cn-beijing.volces.com/api/v3/chat/completions -EMBEDDING_MODEL=text-embedding-v4 NORMOL_MODEL=deepseek-chat HIERARCHY_LLM_MODEL=qwen3.6-flash IMAGE_MODEL=qwen3.5-flash @@ -99,17 +89,8 @@ IMAGE_MODEL_MAX=qwen3.5-flash # File handling defaults SUPPORTED_EXTENSIONS=.doc,.docx,.pdf,.txt,.xls,.xlsx,.csv,.pptx,.jpg,.jpeg,.png,.md MAX_FILE_SIZE=104857600 -MAX_IMAGE_SIZE=10485760 -MIN_CONFIDENCE_THRESHOLD=0.05 -HIGH_IOU_THRESHOLD=0.9 -DEFAULT_EMBEDDING_DIM=1024 -DEFAULT_TOP_K=5 -DEFAULT_BATCH_SIZE=32 -DEFAULT_EPOCHS=3 -DEFAULT_THRESHOLD=0.5 # Required for specific features: webhooks and callbacks -WEBHOOK_SIGNING_SECRET= WEBHOOK_MASTER_KEY= QSTASH_TOKEN= QSTASH_CALLBACK_BASE_URL=https://api.example.com/api/v1 @@ -117,40 +98,14 @@ QSTASH_MAX_RETRIES=5 # QSTASH_CURRENT_SIGNING_KEY= # QSTASH_NEXT_SIGNING_KEY= -# Required for specific features: billing and notifications +# Required for specific features: billing and analytics BILLING_ENABLED=false STRIPE_SECRET_KEY= -STRIPE_PUBLISHABLE_KEY= STRIPE_WEBHOOK_SECRET= -RESEND_API_KEY= -RESEND_FROM_EMAIL=noreply@example.com -RESEND_FROM_NAME=Knowhere -RESEND_MAX_RETRIES=3 -RESEND_RETRY_DELAY=1.0 -RESEND_TEMPLATE_WELCOME= -RESEND_TEMPLATE_PURCHASE_CONFIRMATION= -RESEND_TEMPLATE_JOB_COMPLETION= -RESEND_TEMPLATE_JOB_FAILURE= -RESEND_TEMPLATE_WELCOME_ENABLED=false -RESEND_TEMPLATE_PURCHASE_CONFIRMATION_ENABLED=false -RESEND_TEMPLATE_JOB_COMPLETION_ENABLED=false -RESEND_TEMPLATE_JOB_FAILURE_ENABLED=false MOESIF_APPLICATION_ID= -NEXT_PUBLIC_POSTHOG_KEY= -NEXT_PUBLIC_POSTHOG_HOST=https://app.posthog.com FREE_PLAN_INITIAL_CREDITS=5 FRONTEND_URL=http://localhost:3000 -# Required for specific features: dashboard and auth providers -USERS_VERIFY_TOKEN_SECRET= -USERS_RESET_PASSWORD_TOKEN_SECRET= -GOOGLE_CLIENT_ID= -GOOGLE_CLIENT_SECRET= -GITHUB_CLIENT_ID= -GITHUB_CLIENT_SECRET= -APPLE_CLIENT_ID= -APPLE_CLIENT_SECRET= - # Required for specific features: parsing providers MINERU_API_KEYS= MINERU_URL=https://mineru.net/api/v4 @@ -163,13 +118,6 @@ ILOVEAPI_SECRET_KEY= ILOVEAPI_BASE_URL=https://api.ilovepdf.com/v1 ILOVEAPI_TIMEOUT=120 -# Optional or development-only: compatibility fields kept for retained legacy code paths +# Legacy parser compatibility fields. ALL_DF_COLS=content,path,type,length,keywords,summary,know_id,tokens,connectto,addtime,page_nums -DEFAULT_FOLDERS=Supplementary_Files,Temporary_Files,templates,images,fragments -KB_TERM=KB_DATA -KB_VEC_TERM=KB_VECS -META_PATH=app/core/config/Meta_setting.csv -CONFIG_PATH=app/core/config/config.txt -PATH_IMAGE_PATTERN=.*\.(png|jpe?g|gif)$ -IMG_TBL_PATTERN=\[(?:images|tables)/[^\]]+\] -SPLIT_CHAR=/ +SPLIT_CHAR=--> diff --git a/apps/api/app/api/v1/routes/jobs.py b/apps/api/app/api/v1/routes/jobs.py index 40f0dee33..2cdf7480a 100644 --- a/apps/api/app/api/v1/routes/jobs.py +++ b/apps/api/app/api/v1/routes/jobs.py @@ -362,7 +362,7 @@ async def create_job( # pyright: ignore[reportGeneralTypeIssues] job_type = "kb_management" - # Keep job creation lightweight. The worker reads USERS_DATA_PATH directly. + # Keep job creation lightweight. from shared.services.redis import RedisServiceFactory redis_service = RedisServiceFactory.get_service() diff --git a/apps/api/tests/support/import_environment.py b/apps/api/tests/support/import_environment.py index 8957bdb3e..f06561d08 100644 --- a/apps/api/tests/support/import_environment.py +++ b/apps/api/tests/support/import_environment.py @@ -7,7 +7,6 @@ _REQUIRED_IMPORT_ENVIRONMENT: dict[str, str] = { "DATABASE_URL": "postgresql+asyncpg://user:pass@127.0.0.1:15432/knowhere_test", - "SECRET_KEY": "test-secret-key", "DS_KEY": "test-deepseek-key", "DS_URL": "https://example.com/v1", "S3_BUCKET_NAME": "knowhere-test-bucket", @@ -15,9 +14,6 @@ "S3_SECRET_ACCESS_KEY": "test-secret-key", "S3_TEMP_PATH": "/tmp/knowhere-api-tests", "TMP_PATH": "/tmp/knowhere-api-tests", - "FONT_PATH": "/tmp/knowhere-api-tests", - "CHROMEDRIVER_PATH": "/tmp/knowhere-api-tests/chromedriver", - "USERS_DATA_PATH": "/tmp/knowhere-api-tests/users", } diff --git a/apps/worker/.env.example b/apps/worker/.env.example index 018e5660c..1288b27dc 100644 --- a/apps/worker/.env.example +++ b/apps/worker/.env.example @@ -10,7 +10,7 @@ # Required for specific features: # - alternate object storage callbacks # - webhooks and async callback delivery -# - billing, auth, analytics, and notifications +# - billing and analytics # - alternate parsing providers # # Optional or development-only values can stay empty unless you need the @@ -18,19 +18,12 @@ # Required for local startup: application runtime ENVIRONMENT=development -DEBUG=true LOG_LEVEL=INFO APP_TITLE=Knowhere Worker APP_VERSION=1.0.0 APP_DESCRIPTION=Document parsing and retrieval worker -SECRET_KEY=replace-with-a-long-random-secret -ALGORITHM=HS256 -ACCESS_TOKEN_EXPIRE_MINUTES=10080 API_STANDALONE_MODE_ENABLED=false TMP_PATH=/tmp/knowhere -FONT_PATH=/usr/share/fonts -CHROMEDRIVER_PATH=/usr/bin/chromedriver -USERS_DATA_PATH=/tmp/knowhere-users # Optional or development-only: observability LOGFIRE_TOKEN= @@ -55,7 +48,6 @@ BROKER_POOL_LIMIT=10 # Required for local startup: S3-compatible storage S3_TYPE=s3 S3_BUCKET_NAME=knowhere-uploads -S3_UPLOADS_BUCKET=knowhere-uploads S3_RESULTS_BUCKET=knowhere-results S3_ACCESS_KEY_ID=test S3_SECRET_ACCESS_KEY=test @@ -66,7 +58,6 @@ S3_REGION=us-west-1 S3_USE_SSL=false S3_ADDRESSING_STYLE=path S3_WEBHOOK_AUTH_TOKEN=replace-with-a-shared-secret -SNS_SIGNATURE_VERIFICATION=true # Required for specific features: OSS settings OSS_ENDPOINT= @@ -74,7 +65,6 @@ OSS_EVENT_CALLBACK_KEY= OSS_EVENT_VERIFY_SIGNATURE=true # Required for specific features: webhooks and callbacks -WEBHOOK_SIGNING_SECRET= WEBHOOK_MASTER_KEY= QSTASH_TOKEN= QSTASH_CALLBACK_BASE_URL=https://api.example.com/api/v1 @@ -92,38 +82,16 @@ ALI_API_KEYS= ALI_URL=https://dashscope.aliyuncs.com/compatible-mode/v1 ARK_API_KEY= ARK_URL=https://ark.cn-beijing.volces.com/api/v3/chat/completions -EMBEDDING_MODEL=text-embedding-v4 NORMOL_MODEL=deepseek-chat HIERARCHY_LLM_MODEL=deepseek-chat IMAGE_MODEL=qwen3.5-flash IMAGE_MODEL_MAX=qwen3.5-flash -MIN_CONFIDENCE_THRESHOLD=0.05 -HIGH_IOU_THRESHOLD=0.9 -DEFAULT_EMBEDDING_DIM=1024 -DEFAULT_TOP_K=5 -DEFAULT_BATCH_SIZE=32 -DEFAULT_EPOCHS=3 -DEFAULT_THRESHOLD=0.5 -# Required for specific features: billing, auth, and notifications +# Required for specific features: billing and analytics BILLING_ENABLED=false STRIPE_SECRET_KEY= -STRIPE_PUBLISHABLE_KEY= STRIPE_WEBHOOK_SECRET= -GOOGLE_CLIENT_ID= -GOOGLE_CLIENT_SECRET= -GITHUB_CLIENT_ID= -GITHUB_CLIENT_SECRET= -APPLE_CLIENT_ID= -APPLE_CLIENT_SECRET= -USERS_VERIFY_TOKEN_SECRET= -USERS_RESET_PASSWORD_TOKEN_SECRET= -RESEND_API_KEY= -RESEND_FROM_EMAIL=noreply@example.com -RESEND_FROM_NAME=Knowhere MOESIF_APPLICATION_ID= -NEXT_PUBLIC_POSTHOG_KEY= -NEXT_PUBLIC_POSTHOG_HOST=https://app.posthog.com # Required for specific features: parsing providers MINERU_API_KEYS= @@ -140,15 +108,8 @@ ILOVEAPI_TIMEOUT=120 # File handling defaults SUPPORTED_EXTENSIONS=.doc,.docx,.pdf,.txt,.xls,.xlsx,.csv,.pptx,.jpg,.jpeg,.png,.md MAX_FILE_SIZE=104857600 -MAX_IMAGE_SIZE=10485760 -# Optional or development-only: compatibility fields kept for retained legacy code paths +# Legacy parser compatibility fields. ALL_DF_COLS=content,path,type,length,keywords,summary,know_id,tokens,connectto,addtime,page_nums -DEFAULT_FOLDERS=Supplementary_Files,Temporary_Files,templates,images,fragments -KB_TERM=KB_DATA -KB_VEC_TERM=KB_VECS -META_PATH=app/core/config/Meta_setting.csv -CONFIG_PATH=app/core/config/config.txt -PATH_IMAGE_PATTERN=.*\.(png|jpe?g|gif)$ -IMG_TBL_PATTERN=\[(?:images|tables)/[^\]]+\] -SPLIT_CHAR=/ +SPLIT_CHAR=--> + diff --git a/deploy/local-dev/localstack/init/setup-aws-resources.sh b/deploy/local-dev/localstack/init/setup-aws-resources.sh index 002bbf3f3..46351c275 100755 --- a/deploy/local-dev/localstack/init/setup-aws-resources.sh +++ b/deploy/local-dev/localstack/init/setup-aws-resources.sh @@ -128,7 +128,6 @@ echo " S3_ENDPOINT_URL=http://localhost:4566" echo " S3_ACCESS_KEY_ID=test" echo " S3_SECRET_ACCESS_KEY=test" echo " S3_BUCKET_NAME=knowhere-uploads" -echo " S3_UPLOADS_BUCKET=knowhere-uploads" echo " S3_RESULTS_BUCKET=knowhere-results" echo " S3_REGION=us-west-1" echo " S3_USE_SSL=false" diff --git a/packages/shared-python/shared/core/config/ai.py b/packages/shared-python/shared/core/config/ai.py index 2afc09be2..2b25b5571 100644 --- a/packages/shared-python/shared/core/config/ai.py +++ b/packages/shared-python/shared/core/config/ai.py @@ -14,7 +14,6 @@ class AIConfig(BaseModel): DS_KEY: str = Field(..., description="DeepSeek API key") DS_URL: str = Field(..., description="DeepSeek API URL") GPT_API_KEY: str = Field(default="", description="OpenAI API key") - EMBEDDING_MODEL: str = Field(default="", description="Embedding model") # Default behavior: text/table summaries use deepseek-chat. Hierarchy parsing # can be overridden independently with HIERARCHY_LLM_MODEL. Existing # environment overrides for NORMOL_MODEL / HIERARCHY_LLM_MODEL / @@ -32,21 +31,25 @@ class AIConfig(BaseModel): description="Image model for image summary, atlas, and OCR flows", ) - # Shared model tuning defaults. - MIN_CONFIDENCE_THRESHOLD: float = Field( - default=0.05, description="Minimum confidence threshold" - ) - HIGH_IOU_THRESHOLD: float = Field(default=0.9, description="High-IoU threshold") - DEFAULT_EMBEDDING_DIM: int = Field( - default=1024, description="Default embedding dimension" - ) - DEFAULT_TOP_K: int = Field(default=5, description="Default top-k value") - DEFAULT_BATCH_SIZE: int = Field(default=32, description="Default batch size") - DEFAULT_EPOCHS: int = Field(default=3, description="Default training epochs") - DEFAULT_THRESHOLD: float = Field(default=0.5, description="Default threshold") + IMAGE_MODEL_MAX: str = Field( + default="qwen3.5-flash", + description="Higher-capability image model for OCR and ask-image Q&A", + ) + + # Runtime LLM controls. + LLM_MOCK_ENABLED: bool = Field( + default=False, + description="Short-circuit all OpenAI-compatible LLM calls and return canned mock responses.", + ) + OPENAI_CLIENT_TIMEOUT: int = Field( + default=300, description="OpenAI-compatible client timeout in seconds" + ) + SUMMARY_LLM_MAX_CONCURRENT: int = Field( + default=8, + description="Max concurrent gevent greenlets for parallel post-heading summary LLM calls -- image/table/text (Dashscope).", + ) # Compatibility fields retained during migration. - DX_KEy: str = Field(default="", description="DX key (compatibility field)") ARK_API_KEY: str = Field( default="", description="ARK API key (compatibility field)" ) @@ -79,50 +82,6 @@ class AIConfig(BaseModel): default=3, description="OpenAI SDK max_retries per token for transient 429s (exponential backoff + jitter).", ) - LLM_MOCK_ENABLED: bool = Field( - default=False, - description="Short-circuit all OpenAI-compatible LLM calls and return canned mock responses.", - ) - OPENAI_CLIENT_TIMEOUT: int = Field( - default=300, description="OpenAI-compatible client timeout in seconds" - ) - # Parallel LLM concurrency limits for gevent workers. - # Higher values reduce wall-clock time but increase burst RPM against the - # LLM provider, which can trigger 429s when multiple pods or jobs run - # concurrently. - HEADING_LLM_MAX_CONCURRENT: int = Field( - default=8, - description="Max concurrent gevent greenlets for parallel heading classification LLM calls (DeepSeek).", - ) - SUMMARY_LLM_MAX_CONCURRENT: int = Field( - default=8, - description="Max concurrent gevent greenlets for parallel post-heading summary LLM calls — image/table/text (Dashscope).", - ) - IMAGE_MODEL_MAX: str = Field( - default="qwen3.5-flash", - description="Higher-capability image model for OCR and ask-image Q&A", - ) - REASON_MODEL: str = Field( - default="", description="Reasoning model (compatibility field)" - ) - IMG_HEADER: str = Field( - default="", description="Image header (compatibility field)" - ) - CONFIG_PATH: str = Field( - default="app/core/config/config.txt", - description="Config path (compatibility field)", - ) - META_PATH: str = Field( - default="app/core/config/Meta_setting.csv", - description="Metadata path (compatibility field)", - ) - IMG_TBL_PATTERN: str = Field( - default="", description="Image table pattern (compatibility field)" - ) - PATH_IMAGE_PATTERN: str = Field( - default="", description="Path image pattern (compatibility field)" - ) - SPLIT_CHAR: str = Field(default="/", description="Path separator") ILOVEAPI_PUBLIC_KEY: str = Field( default="", description="iLoveAPI public key (PPTX-to-PDF)" ) @@ -151,21 +110,8 @@ class AIConfig(BaseModel): default=5, description="Max concurrent in-flight iLoveAPI conversions across all workers. Fail-open to LibreOffice when exceeded.", ) - PROD_URL: str = Field( - default="", description="Production URL (compatibility field)" - ) + SPLIT_CHAR: str = Field(default="/", description="Path separator") ALL_DF_COLS: str = Field( default="content,path,type,length,keywords,summary,know_id,tokens,connectto,addtime,page_nums", description="All dataframe columns (compatibility field)", ) - DEFAULT_FOLDERS: str = Field( - default="Supplementary_Files,Temporary_Files,templates,images,fragments", - description="Default folders (compatibility field)", - ) - KB_TERM: str = Field( - default="KB_DATA", description="Knowledge-base term (compatibility field)" - ) - KB_VEC_TERM: str = Field( - default="KB_VECS", - description="Knowledge-base vector term (compatibility field)", - ) diff --git a/packages/shared-python/shared/core/config/base.py b/packages/shared-python/shared/core/config/base.py index 7b414f6d0..0f098cd4e 100644 --- a/packages/shared-python/shared/core/config/base.py +++ b/packages/shared-python/shared/core/config/base.py @@ -16,7 +16,6 @@ class BaseConfig(BaseSettings): default="", description="Deploy environment (|development|staging|production)", ) - DEBUG: bool = Field(default=False, description="Debug mode") LOG_LEVEL: str = Field(default="INFO", description="Log level") # Application metadata. @@ -36,11 +35,6 @@ class BaseConfig(BaseSettings): ) # Security configuration. - SECRET_KEY: str = Field(..., description="JWT secret key") - ALGORITHM: str = Field(default="HS256", description="JWT signing algorithm") - ACCESS_TOKEN_EXPIRE_MINUTES: int = Field( - default=10080, description="Access-token expiration in minutes" - ) WEBHOOK_MASTER_KEY: str = Field( default="", description="Webhook encryption master key" ) @@ -57,8 +51,6 @@ class BaseConfig(BaseSettings): # Local path configuration. TMP_PATH: str = Field(..., description="Temporary-file path") - FONT_PATH: str = Field(..., description="Font-file path") - CHROMEDRIVER_PATH: str = Field(..., description="ChromeDriver path") @field_validator("ENVIRONMENT") @classmethod @@ -89,8 +81,6 @@ def validate_file_paths(self) -> bool: """Validate required local file paths.""" paths_to_check = { "TMP_PATH": self.TMP_PATH, - "FONT_PATH": self.FONT_PATH, - "CHROMEDRIVER_PATH": self.CHROMEDRIVER_PATH, } for name, path in paths_to_check.items(): diff --git a/packages/shared-python/shared/core/config/billing.py b/packages/shared-python/shared/core/config/billing.py index c0b669534..7aa818cdd 100644 --- a/packages/shared-python/shared/core/config/billing.py +++ b/packages/shared-python/shared/core/config/billing.py @@ -21,91 +21,25 @@ class BillingConfig(BaseSettings): STRIPE_SECRET_KEY: Optional[str] = Field( default=None, description="Stripe secret key" ) - STRIPE_PUBLISHABLE_KEY: Optional[str] = Field( - default=None, description="Stripe publishable key" - ) STRIPE_WEBHOOK_SECRET: Optional[str] = Field( default=None, description="Stripe webhook secret" ) - # Subscription plan configuration. - FREE_PLAN_CREDITS: int = Field( - default=100, description="Monthly credits for the free plan" - ) - PLUS_PLAN_CREDITS: int = Field( - default=1000, description="Monthly credits for the Plus plan" - ) - PRO_PLAN_CREDITS: int = Field( - default=10000, description="Monthly credits for the Pro plan" - ) - # Credits (Micro-Dollar System: $1.00 = 1,000,000 micro-credits) MICRO_DOLLARS_PER_PAGE: int = Field( default=1500, description="Micro dollars per page ($0.0015 = 1500 micros)" ) - LOW_BALANCE_THRESHOLD: int = Field( - default=10_000_000, description="low micro dollars threshold, 10 credits" - ) CREDITS_VALID_DAYS: int = Field( default=365, description="Credit validity period in days" ) - # Subscription prices in cents. - PLUS_PLAN_PRICE: int = Field(default=999, description="Plus plan price in cents") - PRO_PLAN_PRICE: int = Field(default=2999, description="Pro plan price in cents") - - # Webhook configuration. - WEBHOOK_SIGNING_SECRET: str = Field(default="default_webhook_secret") - - # Resend email configuration. - RESEND_API_KEY: str = Field(default="") - RESEND_FROM_EMAIL: str = Field( - default="noreply@knowhere.ai", description="Sender email address" - ) - RESEND_FROM_NAME: str = Field( - default="Knowhere AI", description="Sender display name" - ) - RESEND_MAX_RETRIES: int = Field(default=3, description="Maximum retry count") - RESEND_RETRY_DELAY: float = Field(default=1.0, description="Retry delay in seconds") - # Resend template identifiers from the Resend dashboard. - RESEND_TEMPLATE_WELCOME: Optional[str] = Field( - default=None, description="Welcome email template ID" - ) - RESEND_TEMPLATE_PURCHASE_CONFIRMATION: Optional[str] = Field( - default=None, description="Purchase-confirmation email template ID" - ) - RESEND_TEMPLATE_JOB_COMPLETION: Optional[str] = Field( - default=None, description="Job-completion email template ID" - ) - RESEND_TEMPLATE_JOB_FAILURE: Optional[str] = Field( - default=None, description="Job-failure email template ID" - ) - # Resend template feature toggles. - RESEND_TEMPLATE_WELCOME_ENABLED: bool = Field( - default=False, description="Enable the welcome email template" - ) - RESEND_TEMPLATE_PURCHASE_CONFIRMATION_ENABLED: bool = Field( - default=False, description="Enable the purchase-confirmation email template" - ) - RESEND_TEMPLATE_JOB_COMPLETION_ENABLED: bool = Field( - default=False, description="Enable the job-completion email template" - ) - RESEND_TEMPLATE_JOB_FAILURE_ENABLED: bool = Field( - default=False, description="Enable the job-failure email template" - ) - # Moesif configuration. MOESIF_APPLICATION_ID: str = Field(default="") - # PostHog configuration. - NEXT_PUBLIC_POSTHOG_KEY: str = Field(default="") - NEXT_PUBLIC_POSTHOG_HOST: str = Field(default="https://app.posthog.com") - # Subscription defaults. FREE_PLAN_INITIAL_CREDITS: int = Field(default=5) # S3 result-bucket configuration. - S3_UPLOADS_BUCKET: str = Field(default="") S3_RESULTS_BUCKET: str = Field(default="") # Frontend callback URL used by Stripe Checkout. @@ -119,7 +53,6 @@ def validate_billing_config(self) -> bool: if self.STRIPE_SECRET_KEY: required_stripe_fields = [ self.STRIPE_SECRET_KEY, - self.STRIPE_PUBLISHABLE_KEY, self.STRIPE_WEBHOOK_SECRET, ] return all(field for field in required_stripe_fields) diff --git a/packages/shared-python/shared/core/config/storage.py b/packages/shared-python/shared/core/config/storage.py index b660ac385..a7237481f 100644 --- a/packages/shared-python/shared/core/config/storage.py +++ b/packages/shared-python/shared/core/config/storage.py @@ -6,11 +6,10 @@ import boto3 from botocore.client import BaseClient from botocore.config import Config -from pydantic import BaseModel, Field, model_validator +from pydantic import BaseModel, Field from shared.core.exceptions.domain_exceptions import ( DependencyMissingException, - SystemSettingInvalidException, SystemSettingMissingException, ) @@ -56,49 +55,15 @@ class StorageConfig(BaseModel): MAX_FILE_SIZE: int = Field( default=104857600, description="Maximum file size in bytes" ) - MAX_IMAGE_SIZE: int = Field( - default=10485760, description="Maximum image size in bytes" - ) SUPPORTED_EXTENSIONS: str = Field( default=".doc,.docx,.pdf,.txt,.xls,.xlsx,.pptx,.jpg,.jpeg,.png,.md", description="Supported file extensions", ) - # Shared user-data directory for API and worker processes. - USERS_DATA_PATH: str = Field( - ..., description="Absolute path to the shared user-data directory" - ) - - @model_validator(mode="after") - def _validate_users_data_path(self): - """Validate the USERS_DATA_PATH setting.""" - if not self.USERS_DATA_PATH: - raise SystemSettingMissingException( - internal_message="USERS_DATA_PATH must be configured, cannot be empty" - ) - - # Require an absolute path. - if not os.path.isabs(self.USERS_DATA_PATH): - raise SystemSettingInvalidException( - internal_message=f"USERS_DATA_PATH must be an absolute path, current value: {self.USERS_DATA_PATH}" - ) - - # Only check writeability when the directory already exists. - if os.path.exists(self.USERS_DATA_PATH): - if not os.access(self.USERS_DATA_PATH, os.W_OK): - raise SystemSettingInvalidException( - internal_message=f"USERS_DATA_PATH directory is not writable: {self.USERS_DATA_PATH}" - ) - - return self - # S3 event-notification configuration. S3_WEBHOOK_AUTH_TOKEN: str = Field( default="", description="MinIO webhook authentication token" ) - SNS_SIGNATURE_VERIFICATION: bool = Field( - default=True, description="Verify SNS signatures" - ) # OSS event-notification configuration. OSS_EVENT_CALLBACK_KEY: str = Field( diff --git a/packages/shared-python/shared/testing/contract_runtime.py b/packages/shared-python/shared/testing/contract_runtime.py index 8c816ce79..f596f402a 100644 --- a/packages/shared-python/shared/testing/contract_runtime.py +++ b/packages/shared-python/shared/testing/contract_runtime.py @@ -103,12 +103,7 @@ def _ensure_import_paths() -> None: def _ensure_test_directories() -> None: - users_data_path: Path = _TEST_TMP_ROOT / "users" - chromedriver_path: Path = _TEST_TMP_ROOT / "chromedriver" - _TEST_TMP_ROOT.mkdir(parents=True, exist_ok=True) - users_data_path.mkdir(parents=True, exist_ok=True) - chromedriver_path.touch(exist_ok=True) def _reset_contract_storage_state(database_url: str) -> None: @@ -303,8 +298,6 @@ def configure_contract_environment( environment: dict[str, str] = { "ENVIRONMENT": "development", - "DEBUG": "true", - "SECRET_KEY": "test-secret-key", "WEBHOOK_MASTER_KEY": CONTRACT_WEBHOOK_MASTER_KEY, "DATABASE_URL": database_url, "DB_SSL_MODE": "disable", @@ -316,22 +309,17 @@ def configure_contract_environment( f"redis://{CONTRACT_REDIS_HOST}:{CONTRACT_REDIS_PORT}/{CONTRACT_REDIS_DATABASE}" ), "TMP_PATH": str(_TEST_TMP_ROOT), - "FONT_PATH": str(_TEST_TMP_ROOT), - "CHROMEDRIVER_PATH": str(_TEST_TMP_ROOT / "chromedriver"), - "USERS_DATA_PATH": str(_TEST_TMP_ROOT / "users"), "S3_BUCKET_NAME": "knowhere-test-bucket", "S3_ACCESS_KEY_ID": "test-access-key", "S3_SECRET_ACCESS_KEY": "test-secret-key", "S3_TEMP_PATH": str(_TEST_TMP_ROOT), "S3_ENDPOINT_URL": "http://127.0.0.1:4566", "S3_PRIVATE_DOMAIN": "http://127.0.0.1:4566", - "S3_UPLOADS_BUCKET": "knowhere-test-uploads", "S3_RESULTS_BUCKET": "knowhere-test-results", "S3_REGION": "us-west-1", "S3_USE_SSL": "false", "S3_ADDRESSING_STYLE": "path", "STRIPE_SECRET_KEY": "sk_test_contract_secret", - "STRIPE_PUBLISHABLE_KEY": "pk_test_contract_publishable", "STRIPE_WEBHOOK_SECRET": "whsec_contract_test_secret", "DS_KEY": "test-deepseek-key", "DS_URL": "https://example.com/v1", From 1609907fe79a43f6a2b6c993def8397a16f743c7 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Thu, 7 May 2026 00:27:24 +0800 Subject: [PATCH 28/32] refactor: update README and scripts for local development setup; enable billing and rate limiting by default --- README.md | 38 ++-- apps/api/.env.example | 4 +- apps/api/alembic/env.py | 4 +- apps/api/scripts/bootstrap_local_dev.py | 63 ------ apps/api/scripts/init_user.py | 114 ++++++----- .../scripts/local_dev_bootstrap_service.py | 182 ------------------ .../tests/contract/test_api_key_contract.py | 9 +- .../tests/contract/test_billing_contract.py | 2 +- deploy/local-dev/README.md | 30 --- deploy/local-dev/start-dev.sh | 82 +------- .../shared/models/database/job.py | 6 +- .../shared/models/database/job_result.py | 7 +- .../models/database/job_state_audit_log.py | 5 +- .../models/database/job_state_history.py | 5 +- .../shared/models/database/webhook.py | 7 +- .../shared/models/database/webhook_log.py | 6 +- .../shared/testing/contract_runtime.py | 117 ++++++++--- 17 files changed, 209 insertions(+), 472 deletions(-) delete mode 100644 apps/api/scripts/bootstrap_local_dev.py delete mode 100644 apps/api/scripts/local_dev_bootstrap_service.py diff --git a/README.md b/README.md index 42d4ddc47..f0cfa69ae 100644 --- a/README.md +++ b/README.md @@ -62,25 +62,15 @@ cp apps/worker/.env.example apps/worker/.env - `DS_KEY` - any optional LLM, billing, or webhook providers you want to enable -The example files default to the open-source/self-hosted behavior: +These settings control the local startup mode: - `API_STANDALONE_MODE_ENABLED=false` for the combined dashboard + API flow, where the dashboard initializes Better Auth tables before API migrations. -- `BILLING_ENABLED=false`, so Stripe and credit deduction are not required. -- `RATE_LIMIT_ENABLED=false` for local/self-hosted convenience; set it to - `true` when you want API rate limits enforced. +- `BILLING_ENABLED` controls Stripe and credit deduction. +- `RATE_LIMIT_ENABLED` controls API rate limit enforcement. -For API-only development without the dashboard, set `API_STANDALONE_MODE_ENABLED=true`, -run API migrations, then create an API-only user/key: - -```bash -cd apps/api -uv run --python 3.11 python -m alembic upgrade heads -uv run --python 3.11 python scripts/init_user.py --email you@example.com -``` - -If you plan to use the dashboard, start the combined self-hosted stack and -register through the dashboard instead of using `scripts/init_user.py`. +For API-only development without the dashboard, set +`API_STANDALONE_MODE_ENABLED=true` in `apps/api/.env`. 4. Start the local infrastructure stack: @@ -88,20 +78,26 @@ register through the dashboard instead of using `scripts/init_user.py`. ./deploy/local-dev/start-dev.sh ``` -If you also want the helper to initialize the local API user state, rerun it -with `--init-user`: +5. Start the API and worker in separate terminals: ```bash -./deploy/local-dev/start-dev.sh --init-user +cd apps/api && uv run uvicorn main:app --host 0.0.0.0 --port 5005 --reload +cd apps/worker && uv run python worker.py ``` -5. Start the API and worker in separate terminals: +The API runs migrations during startup. + +For API-only development without the dashboard, create an API-only user/key +after the API service starts: ```bash -cd apps/api && uv run main.py -cd apps/worker && uv run worker.py +cd apps/api +uv run --python 3.11 python scripts/init_user.py --email you@example.com ``` +If you plan to use the dashboard, register through the dashboard instead of +using `scripts/init_user.py`. + ## Quality Checks Run lint checks from the repository root: diff --git a/apps/api/.env.example b/apps/api/.env.example index aa995a867..67716c644 100644 --- a/apps/api/.env.example +++ b/apps/api/.env.example @@ -39,7 +39,7 @@ DB_SSL_MODE=disable # DB_SSL_ROOT_CERT=/path/to/ca-cert.pem # Required for local startup: Redis / Celery -RATE_LIMIT_ENABLED=false +RATE_LIMIT_ENABLED=true REDIS_HOST=localhost REDIS_PORT=6379 REDIS_PASSWORD= @@ -99,7 +99,7 @@ QSTASH_MAX_RETRIES=5 # QSTASH_NEXT_SIGNING_KEY= # Required for specific features: billing and analytics -BILLING_ENABLED=false +BILLING_ENABLED=true STRIPE_SECRET_KEY= STRIPE_WEBHOOK_SECRET= MOESIF_APPLICATION_ID= diff --git a/apps/api/alembic/env.py b/apps/api/alembic/env.py index 5dcfb7ede..7b9e5022c 100644 --- a/apps/api/alembic/env.py +++ b/apps/api/alembic/env.py @@ -136,7 +136,7 @@ def run_with_connection(connection: Connection) -> None: return if isinstance(configured_connection, Engine): - with configured_connection.connect() as connection: + with configured_connection.begin() as connection: run_with_connection(connection) return @@ -149,7 +149,7 @@ def run_with_connection(connection: Connection) -> None: connect_args=ssl_connect_args, ) - with connectable.connect() as connection: + with connectable.begin() as connection: run_with_connection(connection) diff --git a/apps/api/scripts/bootstrap_local_dev.py b/apps/api/scripts/bootstrap_local_dev.py deleted file mode 100644 index 779dee503..000000000 --- a/apps/api/scripts/bootstrap_local_dev.py +++ /dev/null @@ -1,63 +0,0 @@ -from __future__ import annotations - -import argparse -import asyncio -import os -import sys - -sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) - -from scripts.local_dev_bootstrap_service import LocalDevelopmentBootstrapService -from shared.core.database import get_db_context - - -def _build_parser() -> argparse.ArgumentParser: - parser = argparse.ArgumentParser( - description="Bootstrap local-only database state for Knowhere API development.", - ) - parser.add_argument( - "--mode", - choices=("ensure-user-table", "seed", "print-profile"), - required=True, - help="Bootstrap mode to run.", - ) - return parser - - -async def _run(mode: str) -> int: - service = LocalDevelopmentBootstrapService() - - if mode == "print-profile": - _print_profile() - return 0 - - if mode == "ensure-user-table": - await service.ensure_user_table_exists() - print('Ensured local development table: "user"') - return 0 - - async with get_db_context() as session: - await service.seed_local_developer(session) - - print("Ensured local development developer account.") - _print_profile() - return 0 - - -def _print_profile() -> None: - profile = LocalDevelopmentBootstrapService.get_local_developer_profile() - print(f"user_id={profile['user_id']}") - print(f"name={profile['name']}") - print(f"email={profile['email']}") - print(f"tier={profile['tier']}") - print(f"api_key={profile['api_key']}") - - -def main() -> int: - parser = _build_parser() - args = parser.parse_args() - return asyncio.run(_run(args.mode)) - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/apps/api/scripts/init_user.py b/apps/api/scripts/init_user.py index fcb9292f5..de38fe61e 100644 --- a/apps/api/scripts/init_user.py +++ b/apps/api/scripts/init_user.py @@ -3,26 +3,32 @@ import argparse import asyncio import os -import secrets import sys +from typing import TypedDict from uuid import uuid4 -from sqlalchemy import select +from sqlalchemy import select, update from sqlalchemy.ext.asyncio import AsyncSession sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) -from shared.core.billing import MicroDollar -from shared.core.config import settings from shared.core.database import engine, get_db_context from shared.models.database.api_key import APIKey from shared.models.database.user import User from shared.models.database.user_balance import UserBalance +from shared.services.billing.credits_service import CreditsService from shared.services.auth.user_table_bootstrap import ensure_better_auth_user_table -from shared.utils.api_keys import hash_api_key +from shared.utils.api_keys import generate_api_key, hash_api_key, mask_api_key _DEFAULT_API_KEY_NAME: str = "standalone-api-key" -_DEFAULT_USER_TIER: str = "free" +_DEFAULT_USER_TIER: str = "tier_5" + + +class InitializedStandaloneUser(TypedDict): + user_id: str + email: str + api_key_name: str + api_key: str def _build_parser() -> argparse.ArgumentParser: @@ -32,6 +38,11 @@ def _build_parser() -> argparse.ArgumentParser: ), ) parser.add_argument("--email", required=True, help="User email address.") + parser.add_argument( + "--user-id", + default="", + help="Optional user ID for deterministic local or test bootstrap.", + ) parser.add_argument("--name", default="", help="Display name for new users.") parser.add_argument( "--key-name", @@ -41,7 +52,7 @@ def _build_parser() -> argparse.ArgumentParser: parser.add_argument( "--tier", default=_DEFAULT_USER_TIER, - help="Compatibility user tier to store in user_balances.", + help="Compatibility user tier to store in user_balances. Examples: tier_1, tier_2, tier_3, tier_4, tier_5, guest", ) return parser @@ -56,11 +67,13 @@ async def _find_or_create_user( *, email: str, name: str, + requested_user_id: str, ) -> User: normalized_email = email.strip().lower() if not normalized_email: raise ValueError("email must not be empty") + normalized_user_id = requested_user_id.strip() result = await session.execute( select(User).where(User.email == normalized_email).limit(1) ) @@ -68,8 +81,16 @@ async def _find_or_create_user( if user is not None: return user + if normalized_user_id: + existing_user = await session.get(User, normalized_user_id) + if existing_user is not None: + raise ValueError( + "requested user_id already exists for a different email: " + f"user_id={normalized_user_id}" + ) + user = User( - id=f"user_{uuid4().hex[:24]}", + id=normalized_user_id or f"user_{uuid4().hex[:24]}", name=name.strip() or normalized_email, email=normalized_email, ) @@ -78,24 +99,16 @@ async def _find_or_create_user( return user -async def _ensure_user_balance( +async def _initialize_user_credits( session: AsyncSession, *, user_id: str, tier: str, ) -> None: - balance = await session.get(UserBalance, user_id) - if balance is not None: - return - - session.add( - UserBalance( - user_id=user_id, - user_tier=tier, - credits_balance=MicroDollar.from_dollars( - settings.FREE_PLAN_INITIAL_CREDITS - ).amount, - ) + credits_service = CreditsService() + await credits_service.ensure_user_initialized(session, user_id) + await session.execute( + update(UserBalance).where(UserBalance.user_id == user_id).values(user_tier=tier) ) @@ -124,28 +137,18 @@ async def _resolve_key_name( return f"{key_name}-{suffix}" -def _generate_api_key() -> str: - return f"sk_kn_{secrets.token_hex(16)}" - - -def _mask_api_key(api_key: str) -> str: - if len(api_key) < 12: - return api_key - return api_key[:8] + "•" * (len(api_key) - 12) + api_key[-4:] - - async def _create_api_key( session: AsyncSession, *, user_id: str, key_name: str, ) -> str: - api_key = _generate_api_key() + api_key = generate_api_key() session.add( APIKey( user_id=user_id, key_hash=hash_api_key(api_key), - key_mask=_mask_api_key(api_key), + key_mask=mask_api_key(api_key), name=key_name, enabled_modules=["all"], ) @@ -153,23 +156,31 @@ async def _create_api_key( return api_key -async def _run(args: argparse.Namespace) -> int: +async def initialize_standalone_user( + *, + email: str, + user_id: str = "", + name: str = "", + key_name: str = _DEFAULT_API_KEY_NAME, + tier: str = _DEFAULT_USER_TIER, +) -> InitializedStandaloneUser: await _ensure_user_table() async with get_db_context() as session: user = await _find_or_create_user( session, - email=str(args.email), - name=str(args.name), + email=email, + name=name, + requested_user_id=user_id, ) - await _ensure_user_balance( + await _initialize_user_credits( session, user_id=user.id, - tier=str(args.tier).strip() or _DEFAULT_USER_TIER, + tier=tier.strip() or _DEFAULT_USER_TIER, ) key_name = await _resolve_key_name( session, user_id=user.id, - requested_name=str(args.key_name), + requested_name=key_name, ) api_key = await _create_api_key( session, @@ -177,10 +188,27 @@ async def _run(args: argparse.Namespace) -> int: key_name=key_name, ) - print(f"user_id={user.id}") - print(f"email={user.email}") - print(f"api_key_name={key_name}") - print(f"api_key={api_key}") + return { + "user_id": str(user.id), + "email": str(user.email), + "api_key_name": key_name, + "api_key": api_key, + } + + +async def _run(args: argparse.Namespace) -> int: + initialized_user = await initialize_standalone_user( + email=str(args.email), + user_id=str(args.user_id), + name=str(args.name), + key_name=str(args.key_name), + tier=str(args.tier), + ) + + print(f"user_id={initialized_user['user_id']}") + print(f"email={initialized_user['email']}") + print(f"api_key_name={initialized_user['api_key_name']}") + print(f"api_key={initialized_user['api_key']}") return 0 diff --git a/apps/api/scripts/local_dev_bootstrap_service.py b/apps/api/scripts/local_dev_bootstrap_service.py deleted file mode 100644 index b88e7e7fe..000000000 --- a/apps/api/scripts/local_dev_bootstrap_service.py +++ /dev/null @@ -1,182 +0,0 @@ -from __future__ import annotations - -from datetime import datetime, timezone - -from sqlalchemy.ext.asyncio import AsyncSession - -from shared.core.billing import MicroDollar -from shared.core.database import engine -from shared.models.database.api_key import APIKey -from shared.models.database.credits_transaction import CreditsTransaction -from shared.models.database.payment_record import PaymentRecord -from shared.models.database.user import User -from shared.models.database.user_balance import UserBalance -from shared.services.auth.user_table_bootstrap import ensure_better_auth_user_table -from shared.utils.api_keys import hash_api_key - - -class LocalDevelopmentBootstrapService: - """Bootstrap local-only user and billing state for development.""" - - LOCAL_DEV_USER_ID: str = "local-dev-user" - LOCAL_DEV_USER_NAME: str = "Local Development User" - LOCAL_DEV_USER_EMAIL: str = "local-dev-user@knowhere.local" - LOCAL_DEV_TIER: str = "tier_5" - LOCAL_DEV_API_KEY_ID: str = "local-dev-default-api-key" - LOCAL_DEV_API_KEY: str = "sk_local_dev_demo_key_tier5_full_access" - LOCAL_DEV_API_KEY_NAME: str = "local-dev-full-access" - LOCAL_DEV_PAYMENT_RECORD_ID: str = "local-dev-seed-payment-record" - LOCAL_DEV_PAYMENT_INTENT_ID: str = "local-dev-seed-highest-tier" - LOCAL_DEV_CREDITS_TRANSACTION_ID: str = "local-dev-seed-credit-entry" - LOCAL_DEV_CREDITS_BALANCE: int = MicroDollar.from_dollars(2_000).amount - LOCAL_DEV_LIFETIME_BILLING_MICRO: int = MicroDollar.from_dollars(2_000).amount - LOCAL_DEV_PAYMENT_AMOUNT_CENTS: int = 200_000 - LOCAL_DEV_FALLBACK_EMAIL_DOMAIN: str = "knowhere.local" - - async def ensure_user_table_exists(self) -> None: - """Create a dashboard-compatible local `user` table needed by API foreign keys.""" - async with engine.begin() as connection: - await connection.run_sync( - ensure_better_auth_user_table, - fallback_email_domain=self.LOCAL_DEV_FALLBACK_EMAIL_DOMAIN, - ) - - async def seed_local_developer(self, session: AsyncSession) -> None: - """Create or refresh the deterministic local top-tier developer account.""" - await self._upsert_user(session) - await self._upsert_user_balance(session) - await self._upsert_payment_record(session) - await self._upsert_credits_transaction(session) - await self._upsert_api_key(session) - await session.flush() - - @classmethod - def get_local_developer_profile(cls) -> dict[str, str | int]: - """Expose deterministic local developer credentials for local tooling.""" - return { - "user_id": cls.LOCAL_DEV_USER_ID, - "name": cls.LOCAL_DEV_USER_NAME, - "email": cls.LOCAL_DEV_USER_EMAIL, - "tier": cls.LOCAL_DEV_TIER, - "api_key": cls.LOCAL_DEV_API_KEY, - "credits_balance": cls.LOCAL_DEV_CREDITS_BALANCE, - "lifetime_billing_micro": cls.LOCAL_DEV_LIFETIME_BILLING_MICRO, - } - - async def _upsert_user(self, session: AsyncSession) -> None: - user = await session.get(User, self.LOCAL_DEV_USER_ID) - if user is None: - session.add( - User( - id=self.LOCAL_DEV_USER_ID, - name=self.LOCAL_DEV_USER_NAME, - email=self.LOCAL_DEV_USER_EMAIL, - ) - ) - return - - user.name = self.LOCAL_DEV_USER_NAME - user.email = self.LOCAL_DEV_USER_EMAIL - - async def _upsert_user_balance(self, session: AsyncSession) -> None: - balance = await session.get(UserBalance, self.LOCAL_DEV_USER_ID) - if balance is None: - session.add( - UserBalance( - user_id=self.LOCAL_DEV_USER_ID, - user_tier=self.LOCAL_DEV_TIER, - credits_balance=self.LOCAL_DEV_CREDITS_BALANCE, - ) - ) - return - - balance.user_tier = self.LOCAL_DEV_TIER - balance.credits_balance = self.LOCAL_DEV_CREDITS_BALANCE - - async def _upsert_payment_record(self, session: AsyncSession) -> None: - payment = await session.get(PaymentRecord, self.LOCAL_DEV_PAYMENT_RECORD_ID) - if payment is None: - session.add( - PaymentRecord( - id=self.LOCAL_DEV_PAYMENT_RECORD_ID, - payment_intent_id=self.LOCAL_DEV_PAYMENT_INTENT_ID, - user_id=self.LOCAL_DEV_USER_ID, - payment_type="local_dev_seed", - amount_cents=self.LOCAL_DEV_PAYMENT_AMOUNT_CENTS, - currency="USD", - status="succeeded", - credits_amount=self.LOCAL_DEV_LIFETIME_BILLING_MICRO, - processed_at=self._utc_now(), - extra_metadata={"reason": "local_dev_seed"}, - ) - ) - return - - payment.payment_intent_id = self.LOCAL_DEV_PAYMENT_INTENT_ID - payment.user_id = self.LOCAL_DEV_USER_ID - payment.payment_type = "local_dev_seed" - payment.amount_cents = self.LOCAL_DEV_PAYMENT_AMOUNT_CENTS - payment.currency = "USD" - payment.status = "succeeded" - payment.credits_amount = self.LOCAL_DEV_LIFETIME_BILLING_MICRO - payment.processed_at = self._utc_now() - payment.extra_metadata = {"reason": "local_dev_seed"} - - async def _upsert_credits_transaction(self, session: AsyncSession) -> None: - transaction = await session.get( - CreditsTransaction, - self.LOCAL_DEV_CREDITS_TRANSACTION_ID, - ) - if transaction is None: - session.add( - CreditsTransaction( - id=self.LOCAL_DEV_CREDITS_TRANSACTION_ID, - user_id=self.LOCAL_DEV_USER_ID, - credits_amount=self.LOCAL_DEV_CREDITS_BALANCE, - transaction_type="local_dev_seed", - description="Local development seed credits", - transaction_metadata={"reason": "local_dev_seed"}, - ) - ) - return - - transaction.user_id = self.LOCAL_DEV_USER_ID - transaction.credits_amount = self.LOCAL_DEV_CREDITS_BALANCE - transaction.transaction_type = "local_dev_seed" - transaction.description = "Local development seed credits" - transaction.transaction_metadata = {"reason": "local_dev_seed"} - - async def _upsert_api_key(self, session: AsyncSession) -> None: - api_key = await session.get(APIKey, self.LOCAL_DEV_API_KEY_ID) - key_hash = hash_api_key(self.LOCAL_DEV_API_KEY) - key_mask = self._mask_api_key(self.LOCAL_DEV_API_KEY) - - if api_key is None: - session.add( - APIKey( - id=self.LOCAL_DEV_API_KEY_ID, - user_id=self.LOCAL_DEV_USER_ID, - key_hash=key_hash, - key_mask=key_mask, - name=self.LOCAL_DEV_API_KEY_NAME, - enabled_modules=["all"], - ) - ) - return - - api_key.user_id = self.LOCAL_DEV_USER_ID - api_key.key_hash = key_hash - api_key.key_mask = key_mask - api_key.name = self.LOCAL_DEV_API_KEY_NAME - api_key.enabled_modules = ["all"] - api_key.is_active = True - - @staticmethod - def _mask_api_key(api_key: str) -> str: - if len(api_key) < 12: - return api_key - return api_key[:8] + "•" * (len(api_key) - 12) + api_key[-4:] - - @staticmethod - def _utc_now() -> datetime: - return datetime.now(timezone.utc).replace(tzinfo=None) diff --git a/apps/api/tests/contract/test_api_key_contract.py b/apps/api/tests/contract/test_api_key_contract.py index eef3053ba..c6690a5a6 100644 --- a/apps/api/tests/contract/test_api_key_contract.py +++ b/apps/api/tests/contract/test_api_key_contract.py @@ -156,6 +156,7 @@ async def test_should_disable_and_then_reenable_an_api_key_via_the_toggle_route( } async with developer_api_client_factory() as api_client: + developer_authorization = api_client.headers["Authorization"] create_response = await api_client.post("/api/v1/auth/create", json=create_payload) assert create_response.status_code == 200 create_response_json = cast(dict[str, object], create_response.json()) @@ -175,9 +176,7 @@ async def test_should_disable_and_then_reenable_an_api_key_via_the_toggle_route( api_client.headers.update({"Authorization": f"Bearer {raw_api_key}"}) pre_toggle_response = await api_client.get("/api/v1/jobs") - api_client.headers.update( - {"Authorization": "Bearer sk_local_dev_demo_key_tier5_full_access"} - ) + api_client.headers.update({"Authorization": developer_authorization}) disable_response = await api_client.put( f"/api/v1/auth/{created_api_key_id}/toggle" ) @@ -185,9 +184,7 @@ async def test_should_disable_and_then_reenable_an_api_key_via_the_toggle_route( api_client.headers.update({"Authorization": f"Bearer {raw_api_key}"}) disabled_key_response = await api_client.get("/api/v1/jobs") - api_client.headers.update( - {"Authorization": "Bearer sk_local_dev_demo_key_tier5_full_access"} - ) + api_client.headers.update({"Authorization": developer_authorization}) enable_response = await api_client.put( f"/api/v1/auth/{created_api_key_id}/toggle" ) diff --git a/apps/api/tests/contract/test_billing_contract.py b/apps/api/tests/contract/test_billing_contract.py index e7ff6be07..9a9a48469 100644 --- a/apps/api/tests/contract/test_billing_contract.py +++ b/apps/api/tests/contract/test_billing_contract.py @@ -26,7 +26,7 @@ async def test_should_return_the_authenticated_users_initialized_credits_balance response = await api_client.get("/api/v1/billing/credits") assert response.status_code == 200 - assert response.json() == {"credits_balance": 2000.0} + assert response.json() == {"credits_balance": 5.0} @pytest.mark.asyncio diff --git a/deploy/local-dev/README.md b/deploy/local-dev/README.md index e481b8033..74cb39dc8 100644 --- a/deploy/local-dev/README.md +++ b/deploy/local-dev/README.md @@ -17,36 +17,6 @@ cd deploy/local-dev ./start-dev.sh ``` -Or run the helper directly: - -```bash -cd deploy/local-dev -./start-dev.sh -``` - -To initialize the local user/auth state too: - -```bash -cd deploy/local-dev -./start-dev.sh --init-user -``` - -The `--init-user` path is idempotent and can be rerun safely. It now: - -- waits for PostgreSQL, Redis, and LocalStack -- forces the bootstrap connection to the local PostgreSQL DSN even if `apps/api/.env` still points somewhere else -- forces `DB_SSL_MODE=disable` for the local bootstrap path -- ensures the local `user` table matches the dashboard-owned schema needed by API migrations -- runs API Alembic migrations in the local environment -- seeds the deterministic local developer account after the local schema is ready - -Deterministic local developer account: - -- `user_id`: `local-dev-user` -- `email`: `local-dev-user@knowhere.local` -- `tier`: `tier_5` -- `api_key`: `local_dev_demo_key_tier5_full_access` - ## Verify the Local API After you start the API process locally, confirm the service is reachable: diff --git a/deploy/local-dev/start-dev.sh b/deploy/local-dev/start-dev.sh index 9864daf67..f0b9000e7 100755 --- a/deploy/local-dev/start-dev.sh +++ b/deploy/local-dev/start-dev.sh @@ -3,39 +3,24 @@ set -euo pipefail SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" -REPO_ROOT="$(cd "${SCRIPT_DIR}/../.." && pwd)" -API_DIR="${REPO_ROOT}/apps/api" COMPOSE_FILE="${SCRIPT_DIR}/docker-compose.dev.yml" -LOCAL_DEV_DATABASE_URL="postgresql+asyncpg://root:root123@localhost:5432/Knowhere" -LOCAL_DEV_DB_SSL_MODE="disable" -RUN_USER_INIT=0 - log_step() { printf '%s\n' "$1" } -warn() { - printf 'Warning: %s\n' "$1" -} - print_usage() { cat </dev/null 2>&1; then - printf 'uv is required for local bootstrap. Install uv first.\n' >&2 - exit 1 - fi -} - require_docker() { if ! docker info >/dev/null 2>&1; then printf 'Docker is not running. Start Docker first.\n' >&2 @@ -115,54 +93,7 @@ wait_for_localstack() { exit 1 } -prepare_api_env() { - if [[ ! -f "${API_DIR}/.env" ]]; then - cp "${API_DIR}/.env.example" "${API_DIR}/.env" - warn "Created apps/api/.env from .env.example for local development." - fi -} - -configure_local_bootstrap_env() { - export DATABASE_URL="${LOCAL_DEV_DATABASE_URL}" - export DB_SSL_MODE="${LOCAL_DEV_DB_SSL_MODE}" -} - -run_local_bootstrap() { - require_uv - prepare_api_env - configure_local_bootstrap_env - - log_step "Ensuring local development user table..." - ( - cd "${API_DIR}" && - uv run --python 3.11 python scripts/bootstrap_local_dev.py --mode ensure-user-table - ) - - log_step "Running local API migrations..." - ( - cd "${API_DIR}" && - uv run --python 3.11 python -m alembic upgrade heads - ) - - log_step "Seeding deterministic local development user..." - ( - cd "${API_DIR}" && - uv run --python 3.11 python scripts/bootstrap_local_dev.py --mode seed - ) -} - print_summary() { - if [[ "${RUN_USER_INIT}" -eq 1 ]]; then - cat < None: connection.close() +def _recreate_contract_database() -> None: + contract_database_url = get_contract_database_url() + contract_database_name = make_url(contract_database_url).database + + if contract_database_name is None: + raise RuntimeError("Contract database URL does not include a database name.") + + _drop_database( + database_name=contract_database_name, + admin_database_url=_build_admin_sync_database_url(contract_database_url), + ) + _ensure_contract_database_exists() + + def _initialize_contract_database() -> None: contract_database_name: str = make_url(get_contract_database_url()).database or "" @@ -576,28 +596,30 @@ def drop_contract_database( _contract_storage_prepared = False -async def _ensure_contract_user_table() -> None: - _ensure_import_paths() - clear_application_modules() - - try: - bootstrap_module: ModuleType = importlib.import_module( - "scripts.local_dev_bootstrap_service" - ) - bootstrap_service = bootstrap_module.LocalDevelopmentBootstrapService() - await bootstrap_service.ensure_user_table_exists() - finally: - await _dispose_async_database_engine() - clear_application_modules() - - def _run_contract_migrations() -> None: + migration_environment = os.environ.copy() + python_path_entries = [ + str(_API_ROOT), + str(_SHARED_ROOT), + migration_environment.get("PYTHONPATH", ""), + ] + migration_environment.update( + { + "API_STANDALONE_MODE_ENABLED": "true", + "DATABASE_URL": get_contract_database_url(), + "DB_SSL_MODE": "disable", + "PYTHONPATH": os.pathsep.join( + entry for entry in python_path_entries if entry + ), + } + ) + result = subprocess.run( [sys.executable, "-m", "alembic", "upgrade", "heads"], cwd=str(_API_ROOT), capture_output=True, text=True, - env=os.environ.copy(), + env=migration_environment, check=False, ) @@ -608,6 +630,32 @@ def _run_contract_migrations() -> None: ) +def _assert_contract_schema_ready() -> None: + connection = psycopg2.connect(_get_contract_sync_database_url()) + + try: + with connection.cursor() as cursor: + cursor.execute( + """ + SELECT tablename + FROM pg_tables + WHERE schemaname = 'public' + AND tablename IN ('tier_limits', 'api_keys', 'user_balances') + """ + ) + migrated_tables = {row[0] for row in cursor.fetchall()} + finally: + connection.close() + + expected_tables = {"tier_limits", "api_keys", "user_balances"} + missing_tables = expected_tables - migrated_tables + if missing_tables: + raise RuntimeError( + "Contract database migration did not create expected tables: " + f"{', '.join(sorted(missing_tables))}" + ) + + async def _create_contract_engine() -> AsyncEngine: return create_async_engine(get_contract_database_url(), future=True) @@ -618,10 +666,10 @@ async def prepare_contract_storage() -> None: _ensure_import_paths() if not _contract_storage_prepared: - _ensure_contract_database_exists() + _recreate_contract_database() _initialize_contract_database() - await _ensure_contract_user_table() _run_contract_migrations() + _assert_contract_schema_ready() _contract_storage_prepared = True await reset_contract_database() @@ -631,21 +679,30 @@ async def prepare_contract_storage() -> None: async def seed_contract_developer() -> dict[str, str | int]: _ensure_import_paths() - bootstrap_module: ModuleType = importlib.import_module( - "scripts.local_dev_bootstrap_service" + init_user_module: ModuleType = importlib.import_module("scripts.init_user") + initialize_standalone_user = getattr( + init_user_module, + "initialize_standalone_user", ) - bootstrap_service = bootstrap_module.LocalDevelopmentBootstrapService() - engine: AsyncEngine = await _create_contract_engine() - session_factory = async_sessionmaker(bind=engine, expire_on_commit=False) try: - async with session_factory() as session: - await bootstrap_service.seed_local_developer(session) - await session.commit() + initialized_user = await initialize_standalone_user( + email=CONTRACT_DEVELOPER_USER_EMAIL, + user_id=CONTRACT_DEVELOPER_USER_ID, + name=CONTRACT_DEVELOPER_USER_NAME, + key_name=CONTRACT_DEVELOPER_API_KEY_NAME, + tier=CONTRACT_DEVELOPER_USER_TIER, + ) finally: - await engine.dispose() + await _dispose_async_database_engine() - return bootstrap_module.LocalDevelopmentBootstrapService.get_local_developer_profile() + return { + "user_id": initialized_user["user_id"], + "name": CONTRACT_DEVELOPER_USER_NAME, + "email": initialized_user["email"], + "tier": CONTRACT_DEVELOPER_USER_TIER, + "api_key": initialized_user["api_key"], + } async def reset_contract_database() -> None: From f67ad0f0566dacd1c11e78cc2ad147263b73cb90 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Thu, 7 May 2026 00:53:10 +0800 Subject: [PATCH 29/32] refactor: update CodeQL config and remove obsolete test files; streamline type annotations --- .github/codeql/codeql-config.yml | 1 + .../test_api_key_user_cache_contract.py | 126 ----------------- .../contract/test_tier_service_contract.py | 98 ------------- .../contract/test_url_upload_contract.py | 70 ---------- .../shared/models/database/job.py | 6 +- .../shared/models/database/webhook_log.py | 4 +- .../shared/tests/utils/test_api_keys.py | 34 ----- .../tests/utils/test_pinned_outbound_http.py | 132 ------------------ 8 files changed, 5 insertions(+), 466 deletions(-) delete mode 100644 apps/api/tests/contract/test_api_key_user_cache_contract.py delete mode 100644 apps/api/tests/contract/test_tier_service_contract.py delete mode 100644 packages/shared-python/shared/tests/utils/test_api_keys.py delete mode 100644 packages/shared-python/shared/tests/utils/test_pinned_outbound_http.py diff --git a/.github/codeql/codeql-config.yml b/.github/codeql/codeql-config.yml index 2bee2a7b3..705bc0dec 100644 --- a/.github/codeql/codeql-config.yml +++ b/.github/codeql/codeql-config.yml @@ -1,4 +1,5 @@ name: "Knowhere CodeQL config" paths-ignore: + - apps/api/scripts/** - packages/shared-python/shared/utils/api_keys.py diff --git a/apps/api/tests/contract/test_api_key_user_cache_contract.py b/apps/api/tests/contract/test_api_key_user_cache_contract.py deleted file mode 100644 index ca9f498d5..000000000 --- a/apps/api/tests/contract/test_api_key_user_cache_contract.py +++ /dev/null @@ -1,126 +0,0 @@ -from datetime import datetime, timedelta, timezone -from importlib import import_module -from typing import cast - -import pytest - -from tests.support.import_environment import configure_import_environment, ensure_import_paths - -from app.services.auth.api_key_service import APIKeyService as APIKeyServiceType -from shared.services.redis.redis_service import RedisService - - -configure_import_environment() -ensure_import_paths() - - -def get_api_key_service_class() -> type[APIKeyServiceType]: - """Import APIKeyService after test import paths are configured.""" - module = import_module("app.services.auth.api_key_service") - return cast(type[APIKeyServiceType], module.APIKeyService) - - -class FakeRedisService: - def __init__(self) -> None: - self.values: dict[str, object] = {} - self.sets: dict[str, set[str]] = {} - self.ttls: dict[str, int] = {} - - async def get(self, key: str) -> object | None: - return self.values.get(key) - - async def set(self, key: str, value: object, ttl: int | None = None) -> bool: - self.values[key] = value - if ttl is not None: - self.ttls[key] = ttl - return True - - async def delete(self, *keys: str) -> int: - deleted_count = 0 - for key in keys: - cached_value = self.values.pop(key, None) - cached_set = self.sets.pop(key, None) - self.ttls.pop(key, None) - if cached_value is not None or cached_set is not None: - deleted_count += 1 - return deleted_count - - async def sadd(self, key: str, *values: object) -> int: - members = self.sets.setdefault(key, set()) - previous_size = len(members) - members.update(str(value) for value in values) - return len(members) - previous_size - - async def srem(self, key: str, *values: object) -> int: - members = self.sets.setdefault(key, set()) - removed_count = 0 - for value in values: - string_value = str(value) - if string_value in members: - members.remove(string_value) - removed_count += 1 - return removed_count - - async def ttl(self, key: str) -> int: - return self.ttls.get(key, -2) - - async def expire(self, key: str, ttl: int) -> bool: - self.ttls[key] = ttl - return True - - -@pytest.mark.asyncio -async def test_api_key_cache_should_store_user_id_without_tier() -> None: - service = get_api_key_service_class().get_instance() - fake_redis = FakeRedisService() - redis_service = cast(RedisService, fake_redis) - - await service._set_cached_user_id( - redis_service, - api_key_hash="hash-one", - user_id="user-one", - ttl_seconds=7200, - ) - - user_id_key = service._get_user_id_key("hash-one") - user_api_keys_key = service._get_user_api_keys_key("user-one") - - assert await service._get_cached_user_id(redis_service, "hash-one") == "user-one" - assert fake_redis.values[user_id_key] == "user-one" - assert fake_redis.ttls[user_id_key] == 3600 - assert fake_redis.sets[user_api_keys_key] == {"hash-one"} - assert fake_redis.ttls[user_api_keys_key] == 3600 - - -@pytest.mark.asyncio -async def test_api_key_cache_invalidation_should_not_touch_tier_cache() -> None: - service = get_api_key_service_class().get_instance() - fake_redis = FakeRedisService() - redis_service = cast(RedisService, fake_redis) - - await service._set_cached_user_id(redis_service, "hash-one", "user-one", 300) - fake_redis.values["tier:user:user-one"] = "tier_5" - - await service._invalidate_cached_api_key_user_id( - redis_service, - user_id="user-one", - api_key_hash="hash-one", - ) - - assert await service._get_cached_user_id(redis_service, "hash-one") is None - assert fake_redis.values["tier:user:user-one"] == "tier_5" - - -def test_api_key_cache_ttl_should_not_exceed_api_key_expiration() -> None: - service = get_api_key_service_class().get_instance() - expires_at = datetime.now(timezone.utc) + timedelta(seconds=120) - - ttl_seconds = service._resolve_api_key_cache_ttl_seconds(expires_at) - - assert 1 <= ttl_seconds <= 120 - - -def test_api_key_service_should_be_singleton() -> None: - api_key_service_class = get_api_key_service_class() - - assert api_key_service_class.get_instance() is api_key_service_class() diff --git a/apps/api/tests/contract/test_tier_service_contract.py b/apps/api/tests/contract/test_tier_service_contract.py deleted file mode 100644 index 6370e16a5..000000000 --- a/apps/api/tests/contract/test_tier_service_contract.py +++ /dev/null @@ -1,98 +0,0 @@ -from importlib import import_module -from typing import cast - -import pytest - -from tests.support.import_environment import configure_import_environment, ensure_import_paths - -from app.services.rate_limit.tier_service import TierService as TierServiceType -from shared.core.exceptions.domain_exceptions import NotFoundException -from shared.services.redis.redis_service import RedisService - -configure_import_environment() -ensure_import_paths() - - -def get_tier_service_class() -> type[TierServiceType]: - """Import TierService after test import paths are configured.""" - module = import_module("app.services.rate_limit.tier_service") - return cast(type[TierServiceType], module.TierService) - - -def get_not_found_exception_class() -> type[NotFoundException]: - """Import NotFoundException after test import paths are configured.""" - module = import_module("shared.core.exceptions.domain_exceptions") - return cast(type[NotFoundException], module.NotFoundException) - - -class FakeRedisService: - def __init__(self) -> None: - self.values: dict[str, object] = {} - self.ttls: dict[str, int] = {} - - async def get(self, key: str) -> object | None: - return self.values.get(key) - - async def set(self, key: str, value: object, ttl: int | None = None) -> bool: - self.values[key] = value - if ttl is not None: - self.ttls[key] = ttl - return True - - async def delete(self, *keys: str) -> int: - deleted_count = 0 - for key in keys: - cached_value = self.values.pop(key, None) - self.ttls.pop(key, None) - if cached_value is not None: - deleted_count += 1 - return deleted_count - - -@pytest.mark.asyncio -async def test_tier_cache_should_store_user_tier_without_identity_payload() -> None: - tier_service_class = get_tier_service_class() - fake_redis = FakeRedisService() - redis_service = cast(RedisService, fake_redis) - - await tier_service_class._set_cached_tier(redis_service, "user-one", "tier_5") - - cache_key = tier_service_class._get_user_tier_key("user-one") - assert await tier_service_class._get_cached_tier(redis_service, "user-one") == "tier_5" - assert fake_redis.values[cache_key] == "tier_5" - assert fake_redis.ttls[cache_key] == 3600 - - -@pytest.mark.asyncio -async def test_tier_cache_invalidation_should_only_delete_user_tier_key() -> None: - tier_service_class = get_tier_service_class() - fake_redis = FakeRedisService() - redis_service = cast(RedisService, fake_redis) - - await tier_service_class._set_cached_tier(redis_service, "user-one", "tier_5") - fake_redis.values["api-key:user-one"] = "should-stay" - - await redis_service.delete(tier_service_class._get_user_tier_key("user-one")) - - assert await tier_service_class._get_cached_tier(redis_service, "user-one") is None - assert fake_redis.values["api-key:user-one"] == "should-stay" - - -@pytest.mark.asyncio -async def test_get_tier_from_db_should_raise_when_user_tier_is_missing() -> None: - tier_service_class = get_tier_service_class() - not_found_exception_class = get_not_found_exception_class() - - class EmptySession: - async def execute(self, statement: object) -> object: - class EmptyResult: - def scalar_one_or_none(self) -> object | None: - return None - - return EmptyResult() - - with pytest.raises(not_found_exception_class): - await tier_service_class._get_tier_from_db( - cast(object, EmptySession()), - "missing-user", - ) diff --git a/apps/worker/tests/contract/test_url_upload_contract.py b/apps/worker/tests/contract/test_url_upload_contract.py index 6ccd4e945..6a4674403 100644 --- a/apps/worker/tests/contract/test_url_upload_contract.py +++ b/apps/worker/tests/contract/test_url_upload_contract.py @@ -142,73 +142,3 @@ def resolve_public_address( assert job_row["source_type"] == "url" assert job_row["s3_key"] == s3_key - -def test_should_download_a_url_file_through_a_pinned_public_ip( - worker_contract_environment: None, - monkeypatch: MonkeyPatch, - tmp_path: Path, -) -> None: - import app.services.storage.sync_storage_service as sync_storage_service - - source_url = "https://example.test/files/contract-source.pdf" - pinned_ip = "93.184.216.34" - validation_calls: list[tuple[str, str]] = [] - download_calls: list[dict[str, object]] = [] - - def fake_validate_public_http_url_and_resolve_ip( - url: str, - field: str = "url", - ) -> SimpleNamespace: - validation_calls.append((url, field)) - return SimpleNamespace(url=url, validated_ip=pinned_ip) - - def fake_download_pinned_outbound_file( - *, - url: str, - pinned_ip: str, - timeout_seconds: float, - user_agent: str, - temp_dir: str | None = None, - field: str = "source_url", - ) -> SimpleNamespace: - download_calls.append( - { - "url": url, - "pinned_ip": pinned_ip, - "timeout_seconds": timeout_seconds, - "user_agent": user_agent, - "temp_dir": temp_dir, - "field": field, - } - ) - temp_file_path = Path(temp_dir or tmp_path) / "downloaded-contract-source.pdf" - temp_file_path.write_bytes(b"pdf") - return SimpleNamespace(status=200, temp_file_path=str(temp_file_path)) - - monkeypatch.setattr( - sync_storage_service, - "validate_public_http_url_and_resolve_ip", - fake_validate_public_http_url_and_resolve_ip, - ) - monkeypatch.setattr( - sync_storage_service, - "download_pinned_outbound_file", - fake_download_pinned_outbound_file, - ) - monkeypatch.setattr(sync_storage_service.settings, "TMP_PATH", str(tmp_path)) - - downloaded_path = sync_storage_service.download_file_from_url(source_url) - - assert validation_calls == [(source_url, "source_url")] - assert download_calls == [ - { - "url": source_url, - "pinned_ip": pinned_ip, - "timeout_seconds": 300, - "user_agent": "Knowhere-FileDownloader/1.0", - "temp_dir": str(tmp_path), - "field": "source_url", - } - ] - assert downloaded_path == str(tmp_path / "downloaded-contract-source.pdf") - assert Path(downloaded_path).read_bytes() == b"pdf" diff --git a/packages/shared-python/shared/models/database/job.py b/packages/shared-python/shared/models/database/job.py index 1bc2f152e..2613b8722 100644 --- a/packages/shared-python/shared/models/database/job.py +++ b/packages/shared-python/shared/models/database/job.py @@ -27,11 +27,11 @@ from shared.core.database import Base from shared.utils.utc_now import utc_now_naive -from shared.models.database.job_state_audit_log import JobStateAuditLog -from shared.models.database.job_state_history import JobStateHistory -from shared.models.database.webhook_log import WebhookLog if TYPE_CHECKING: + from shared.models.database.job_state_history import JobStateHistory + from shared.models.database.webhook_log import WebhookLog + from shared.models.database.job_state_audit_log import JobStateAuditLog from shared.models.database.job_result import JobResult diff --git a/packages/shared-python/shared/models/database/webhook_log.py b/packages/shared-python/shared/models/database/webhook_log.py index 9bafb2ad3..2beb6dca2 100644 --- a/packages/shared-python/shared/models/database/webhook_log.py +++ b/packages/shared-python/shared/models/database/webhook_log.py @@ -16,11 +16,9 @@ from shared.core.database import Base from shared.utils.utc_now import utc_now_naive -from shared.models.database.webhook import WebhookEvent - if TYPE_CHECKING: from shared.models.database.job import Job - + from shared.models.database.webhook import WebhookEvent class WebhookLog(Base): """Webhook Log Model - Records webhook delivery history.""" diff --git a/packages/shared-python/shared/tests/utils/test_api_keys.py b/packages/shared-python/shared/tests/utils/test_api_keys.py deleted file mode 100644 index c51d924f1..000000000 --- a/packages/shared-python/shared/tests/utils/test_api_keys.py +++ /dev/null @@ -1,34 +0,0 @@ -from shared.utils.api_keys import ( - API_KEY_PREFIX, - generate_api_key, - hash_api_key, - is_api_key_token, - mask_api_key, -) - - -def test_generate_api_key_should_use_api_key_prefix_and_random_secret() -> None: - first_api_key: str = generate_api_key() - second_api_key: str = generate_api_key() - - assert first_api_key.startswith(API_KEY_PREFIX) - assert second_api_key.startswith(API_KEY_PREFIX) - assert first_api_key != second_api_key - assert len(first_api_key) > len(API_KEY_PREFIX) + 32 - - -def test_hash_api_key_should_return_deterministic_sha256_lookup_hash() -> None: - api_key: str = "sk_contract_test_secret" - - assert hash_api_key(api_key) == hash_api_key(api_key) - assert len(hash_api_key(api_key)) == 64 - - -def test_mask_api_key_should_hide_middle_characters() -> None: - assert mask_api_key("sk_1234567890abcdef") == "sk_12345•••••••cdef" - - -def test_is_api_key_token_should_match_only_api_key_prefix() -> None: - assert is_api_key_token("sk_test") is True - assert is_api_key_token("jwt_test") is False - assert is_api_key_token(None) is False diff --git a/packages/shared-python/shared/tests/utils/test_pinned_outbound_http.py b/packages/shared-python/shared/tests/utils/test_pinned_outbound_http.py deleted file mode 100644 index e74d02c02..000000000 --- a/packages/shared-python/shared/tests/utils/test_pinned_outbound_http.py +++ /dev/null @@ -1,132 +0,0 @@ -from __future__ import annotations - -from pathlib import Path -from typing import Any - -import pytest - -from shared.core.exceptions.domain_exceptions import ValidationException -from shared.utils import pinned_outbound_http - - -class _RedirectResponse: - status = 302 - - def stream(self, chunk_size: int) -> list[bytes]: - return [] - - def release_conn(self) -> None: - return None - - def close(self) -> None: - return None - - -class _RedirectConnectionPool: - def __init__(self, *args: object, **kwargs: object) -> None: - self.args = args - self.kwargs = kwargs - - def urlopen(self, *args: object, **kwargs: object) -> _RedirectResponse: - assert kwargs["redirect"] is False - return _RedirectResponse() - - -class _SuccessResponse: - status = 200 - - def __init__(self) -> None: - self.is_released = False - self.is_closed = False - - def stream(self, chunk_size: int) -> list[bytes]: - return [b"pdf", b""] - - def release_conn(self) -> None: - self.is_released = True - - def close(self) -> None: - self.is_closed = True - - -class _SuccessConnectionPool: - calls: list[dict[str, Any]] = [] - - def __init__(self, *args: object, **kwargs: object) -> None: - self.args = args - self.kwargs = kwargs - - def urlopen(self, *args: object, **kwargs: object) -> _SuccessResponse: - self.calls.append( - { - "init_args": self.args, - "init_kwargs": self.kwargs, - "urlopen_args": args, - "urlopen_kwargs": kwargs, - } - ) - return _SuccessResponse() - - -def test_should_block_redirect_responses_and_remove_partial_download( - monkeypatch: pytest.MonkeyPatch, - tmp_path: Path, -) -> None: - monkeypatch.setattr( - pinned_outbound_http, - "PinnedHTTPConnectionPool", - _RedirectConnectionPool, - ) - - with pytest.raises(ValidationException): - pinned_outbound_http.download_pinned_outbound_file( - url="http://example.test/file.pdf", - pinned_ip="93.184.216.34", - timeout_seconds=300, - user_agent="Knowhere-FileDownloader/1.0", - temp_dir=str(tmp_path), - ) - - assert list(tmp_path.iterdir()) == [] - - -def test_should_request_public_url_through_the_pinned_http_pool( - monkeypatch: pytest.MonkeyPatch, - tmp_path: Path, -) -> None: - _SuccessConnectionPool.calls = [] - monkeypatch.setattr( - pinned_outbound_http, - "PinnedHTTPConnectionPool", - _SuccessConnectionPool, - ) - - result = pinned_outbound_http.download_pinned_outbound_file( - url="http://example.test:8080/files/source.pdf?download=1", - pinned_ip="93.184.216.34", - timeout_seconds=300, - user_agent="Knowhere-FileDownloader/1.0", - temp_dir=str(tmp_path), - ) - - assert Path(result.temp_file_path).read_bytes() == b"pdf" - - call = _SuccessConnectionPool.calls[0] - init_kwargs = call["init_kwargs"] - urlopen_kwargs = call["urlopen_kwargs"] - retry_config = init_kwargs["retries"] - timeout = urlopen_kwargs["timeout"] - - assert call["init_args"] == ("example.test", 8080) - assert init_kwargs["pinned_ip"] == "93.184.216.34" - assert retry_config.total == 0 - assert retry_config.redirect == 0 - assert call["urlopen_args"] == ("GET", "/files/source.pdf?download=1") - assert timeout.connect_timeout == 300 - assert timeout.read_timeout == 300 - assert urlopen_kwargs["preload_content"] is False - assert urlopen_kwargs["redirect"] is False - assert urlopen_kwargs["headers"] == { - "User-Agent": "Knowhere-FileDownloader/1.0", - "Host": "example.test:8080", - } From abcc9dd749035c1c3af8c71861f42b1a7d6fee1e Mon Sep 17 00:00:00 2001 From: suguanYang Date: Thu, 7 May 2026 01:06:24 +0800 Subject: [PATCH 30/32] refactor: update README and environment files for API standalone mode; remove Moesif middleware and clean up documentation --- README.md | 16 +- apps/api/.env.example | 2 +- apps/api/app/middleware/moesif_middleware.py | 258 ------------------ apps/api/main.py | 6 - apps/worker/.env.example | 1 - deploy/local-dev/start-dev.sh | 2 - docs/external-services.md | 6 +- .../shared-python/shared/utils/api_keys.py | 1 - 8 files changed, 5 insertions(+), 287 deletions(-) delete mode 100644 apps/api/app/middleware/moesif_middleware.py diff --git a/README.md b/README.md index f0cfa69ae..7335e9e1f 100644 --- a/README.md +++ b/README.md @@ -62,16 +62,6 @@ cp apps/worker/.env.example apps/worker/.env - `DS_KEY` - any optional LLM, billing, or webhook providers you want to enable -These settings control the local startup mode: - -- `API_STANDALONE_MODE_ENABLED=false` for the combined dashboard + API flow, where - the dashboard initializes Better Auth tables before API migrations. -- `BILLING_ENABLED` controls Stripe and credit deduction. -- `RATE_LIMIT_ENABLED` controls API rate limit enforcement. - -For API-only development without the dashboard, set -`API_STANDALONE_MODE_ENABLED=true` in `apps/api/.env`. - 4. Start the local infrastructure stack: ```bash @@ -81,8 +71,8 @@ For API-only development without the dashboard, set 5. Start the API and worker in separate terminals: ```bash -cd apps/api && uv run uvicorn main:app --host 0.0.0.0 --port 5005 --reload -cd apps/worker && uv run python worker.py +cd apps/api && uv run main.py +cd apps/worker && uv run worker.py ``` The API runs migrations during startup. @@ -92,7 +82,7 @@ after the API service starts: ```bash cd apps/api -uv run --python 3.11 python scripts/init_user.py --email you@example.com +uv run scripts/init_user.py --email you@example.com ``` If you plan to use the dashboard, register through the dashboard instead of diff --git a/apps/api/.env.example b/apps/api/.env.example index 67716c644..54c9a9709 100644 --- a/apps/api/.env.example +++ b/apps/api/.env.example @@ -24,7 +24,7 @@ APP_TITLE=Knowhere API APP_VERSION=1.0.0 APP_DESCRIPTION=Document ingestion, retrieval, and MCP backend INTERNAL_DASHBOARD_ENDPOINT=http://localhost:3000 -API_STANDALONE_MODE_ENABLED=false +API_STANDALONE_MODE_ENABLED=true TMP_PATH=/tmp/knowhere # Optional or development-only: observability and local dashboard wiring diff --git a/apps/api/app/middleware/moesif_middleware.py b/apps/api/app/middleware/moesif_middleware.py deleted file mode 100644 index e714d8967..000000000 --- a/apps/api/app/middleware/moesif_middleware.py +++ /dev/null @@ -1,258 +0,0 @@ -""" -Moesif API monitoring middleware. -""" - -import json -import time -from typing import Any, Dict, Optional - -from fastapi import Request, Response -from loguru import logger -from starlette.middleware.base import BaseHTTPMiddleware -from starlette.types import ASGIApp - -from shared.core.config import settings - - -class MoesifMiddleware(BaseHTTPMiddleware): - """Send request and response telemetry to Moesif.""" - - def __init__(self, app: ASGIApp, moesif_application_id: Optional[str] = None): - super().__init__(app) - self.moesif_application_id = ( - moesif_application_id or settings.MOESIF_APPLICATION_ID - ) - self.moesif_client: object | None = None - - if self.moesif_application_id: - try: - from moesifapi.configuration import Configuration - from moesifapi.moesif_api_client import MoesifAPIClient - - configuration = Configuration() - setattr(configuration, "api_key", self.moesif_application_id) - - self.moesif_client = MoesifAPIClient(configuration) - logger.info("Initialized the Moesif client") - - except ImportError: - logger.warning("Moesif SDK is not installed; skipping API monitoring") - except Exception as e: - logger.error(f"Failed to initialize the Moesif client: {e}") - - async def dispatch(self, request: Request, call_next): - """Capture request/response telemetry for one request.""" - start_time = time.time() - - # Capture request information. - request_data = await self._extract_request_data(request) - - # Run the downstream handler. - response = await call_next(request) - - # Measure total processing time. - process_time = time.time() - start_time - - # Capture response information. - response_data = self._extract_response_data(response, process_time) - - # Send the event to Moesif when configured. - if self.moesif_client: - await self._send_to_moesif(request_data, response_data) - - return response - - async def _extract_request_data(self, request: Request) -> Dict[str, Any]: - """Extract request metadata for Moesif.""" - try: - # Read the request body when the method usually carries one. - body = None - if request.method in ["POST", "PUT", "PATCH"]: - try: - body = await request.body() - if body: - # Prefer decoded JSON when possible. - try: - body = json.loads(body.decode()) - except (json.JSONDecodeError, UnicodeDecodeError): - # Fall back to a decoded string when the body is not JSON. - body = body.decode("utf-8", errors="ignore") - except Exception: - body = None - - # Read query parameters. - query_params = dict(request.query_params) - - # Resolve the user identifier from auth or forwarded headers. - user_id = await self._get_user_id(request) - - # Read the forwarded session token when present. - session_token = request.headers.get("x-session-token") - - return { - "time": int(time.time() * 1000), # Millisecond timestamp. - "uri": str(request.url), - "verb": request.method, - "headers": dict(request.headers), - "api_version": request.headers.get("x-api-version", "1.0"), - "ip_address": request.client.host if request.client else None, - "user_id": user_id, - "session_token": session_token, - "body": body, - "query_params": query_params, - } - - except Exception as e: - logger.error(f"Failed to extract request data: {e}") - return {} - - def _extract_response_data( - self, response: Response, process_time: float - ) -> Dict[str, Any]: - """Extract response metadata for Moesif.""" - try: - return { - "time": int(time.time() * 1000), - "status": response.status_code, - "headers": dict(response.headers), - "body": None, # Response bodies are usually too large to store here. - "transfer_encoding": response.headers.get("transfer-encoding"), - "content_length": response.headers.get("content-length"), - "process_time_ms": round(process_time * 1000, 2), - } - - except Exception as e: - logger.error(f"Failed to extract response data: {e}") - return {} - - async def _get_user_id(self, request: Request) -> Optional[str]: - """Resolve the user identifier for telemetry.""" - try: - # Check the Authorization header first. - auth_header = request.headers.get("authorization") - if auth_header: - if auth_header.startswith("Bearer "): - # JWT token. - auth_header[7:] - # TODO: Parse the JWT and extract the user ID. - elif auth_header.startswith("ApiKey "): - # API key. - auth_header[7:] - # TODO: Resolve the user ID from the API key in the database. - - # Fall back to X-User-ID when the frontend provides it. - return request.headers.get("x-user-id") - - except Exception as e: - logger.error(f"Failed to resolve user ID: {e}") - return None - - async def _send_to_moesif( - self, request_data: Dict[str, Any], response_data: Dict[str, Any] - ): - """Send one event to Moesif.""" - try: - if not self.moesif_client: - return - - # Build the Moesif event payload. - event = { - "request": request_data, - "response": response_data, - "user_id": request_data.get("user_id"), - "session_token": request_data.get("session_token"), - "tags": self._get_event_tags(request_data, response_data), - "metadata": self._get_event_metadata(request_data, response_data), - } - - # Send asynchronously without blocking the request. - import asyncio - - asyncio.create_task(self._send_event_async(event)) - - except Exception as e: - logger.error(f"Failed to send Moesif event: {e}") - - async def _send_event_async(self, event: Dict[str, Any]): - """Send one event to Moesif asynchronously.""" - try: - # Moesif's Python SDK is synchronous, so send it in a thread pool. - import asyncio - import concurrent.futures - - def send_sync(): - try: - client = self.moesif_client - if client is None: - return - - # Use whichever send method the installed client exposes. - create_event = getattr(client, "create_event", None) - create_events = getattr(client, "create_events", None) - if callable(create_event): - create_event(event) - elif callable(create_events): - create_events([event]) - else: - logger.warning( - "The Moesif client does not expose create_event or create_events" - ) - except Exception as e: - logger.error(f"Synchronous Moesif send failed: {e}") - - # Run the synchronous client call in a thread pool. - loop = asyncio.get_event_loop() - with concurrent.futures.ThreadPoolExecutor() as executor: - await loop.run_in_executor(executor, send_sync) - - except Exception as e: - logger.error(f"Async Moesif send failed: {e}") - - def _get_event_tags( - self, request_data: Dict[str, Any], response_data: Dict[str, Any] - ) -> Dict[str, str]: - """Build event tags for Moesif analytics.""" - tags = {} - - # Tag by feature area. - uri = request_data.get("uri", "") - if "/billing" in uri: - tags["feature"] = "billing" - elif "/auth" in uri: - tags["feature"] = "authentication" - - # Tag by response status family. - status = response_data.get("status", 200) - if 200 <= status < 300: - tags["status"] = "success" - elif 400 <= status < 500: - tags["status"] = "client_error" - elif 500 <= status < 600: - tags["status"] = "server_error" - - return tags - - def _get_event_metadata( - self, request_data: Dict[str, Any], response_data: Dict[str, Any] - ) -> Dict[str, Any]: - """Build event metadata for Moesif analytics.""" - metadata = {} - - # Include total processing time. - process_time = response_data.get("process_time_ms", 0) - metadata["process_time_ms"] = process_time - - # Include request size when known. - body = request_data.get("body") - if body: - if isinstance(body, str): - metadata["request_size_bytes"] = len(body.encode()) - elif isinstance(body, dict): - metadata["request_size_bytes"] = len(json.dumps(body).encode()) - - # Include response size when known. - content_length = response_data.get("content_length") - if content_length: - metadata["response_size_bytes"] = int(content_length) - - return metadata diff --git a/apps/api/main.py b/apps/api/main.py index c77dd5d6a..ed2c2ae6d 100644 --- a/apps/api/main.py +++ b/apps/api/main.py @@ -137,12 +137,6 @@ def create_app() -> FastAPI: setup_cors(app) app.add_middleware(LoggingMiddleware) - # Moesif API monitoring middleware — disabled (broken SDK client, adds latency + log noise) - # app.add_middleware(MoesifMiddleware) - - # Add API Key authentication middleware - # app.add_middleware(api_key_auth_middleware) - @app.get("/", tags=["Root"]) async def read_root(): return {"message": f"Welcome to {app.title} - Knowledge Base API Service!"} diff --git a/apps/worker/.env.example b/apps/worker/.env.example index 1288b27dc..49bb7aae4 100644 --- a/apps/worker/.env.example +++ b/apps/worker/.env.example @@ -22,7 +22,6 @@ LOG_LEVEL=INFO APP_TITLE=Knowhere Worker APP_VERSION=1.0.0 APP_DESCRIPTION=Document parsing and retrieval worker -API_STANDALONE_MODE_ENABLED=false TMP_PATH=/tmp/knowhere # Optional or development-only: observability diff --git a/deploy/local-dev/start-dev.sh b/deploy/local-dev/start-dev.sh index f0b9000e7..6224235fe 100755 --- a/deploy/local-dev/start-dev.sh +++ b/deploy/local-dev/start-dev.sh @@ -106,8 +106,6 @@ Service endpoints: Next steps: 1. Start the API: cd apps/api && uv run uvicorn main:app --host 0.0.0.0 --port 5005 --reload 2. Start the worker: cd apps/worker && uv run python worker.py - 3. For API-only development, create a dev user after the API starts: - cd apps/api && uv run --python 3.11 python scripts/init_user.py --email you@example.com Stop services: ${SCRIPT_DIR}/stop-dev.sh diff --git a/docs/external-services.md b/docs/external-services.md index 9cb8b0033..32ecb405f 100644 --- a/docs/external-services.md +++ b/docs/external-services.md @@ -6,7 +6,7 @@ references are helpful. ## Required For Local Startup -An external contributor needs these dependencies to run the retained backend +Needs these dependencies to run the backend surface locally: - PostgreSQL for the application database @@ -27,16 +27,12 @@ LocalStack so the default `env.example` files can use a coherent local baseline. required only if you want queued outbound webhook delivery - Stripe: required only if you enable billing and checkout flows -- Resend: - required only if you enable email notifications - OAuth provider credentials: required only if you run dashboard-linked auth flows ## Optional Observability And Analytics - Logfire for distributed tracing export -- Moesif for API analytics -- PostHog for product analytics These integrations are intentionally optional. Leaving them empty should not block a local backend bootstrap. diff --git a/packages/shared-python/shared/utils/api_keys.py b/packages/shared-python/shared/utils/api_keys.py index e119eead3..1fac1befc 100644 --- a/packages/shared-python/shared/utils/api_keys.py +++ b/packages/shared-python/shared/utils/api_keys.py @@ -8,7 +8,6 @@ API_KEY_RANDOM_BYTES: int = 32 -# TODO, use an alphanumeric api key def generate_api_key() -> str: """Generate a new plaintext API key with cryptographic randomness.""" return f"{API_KEY_PREFIX}{token_urlsafe(API_KEY_RANDOM_BYTES)}" From ac5b524d19b35c7dab3bec153d9872e5900ddce2 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Thu, 7 May 2026 11:23:22 +0800 Subject: [PATCH 31/32] refactor: replace outbound URL validation with HTTP URL validation across multiple files; remove obsolete validator --- apps/api/app/api/v1/routes/jobs.py | 8 +- apps/api/app/api/v1/routes/s3_events.py | 8 +- .../contract/test_job_creation_contract.py | 138 ++++++- .../services/storage/sync_storage_service.py | 15 +- .../contract/test_url_upload_contract.py | 1 - .../test_webhook_recovery_contract.py | 4 +- .../shared/models/database/job.py | 4 +- .../services/storage/file_upload_service.py | 15 +- .../shared/services/webhook/dispatcher.py | 10 +- .../services/webhook/qstash_publisher.py | 8 +- .../shared/testing/contract_runtime.py | 2 +- .../shared/utils/outbound_url_validator.py | 91 ----- .../shared/utils/url_file_type.py | 54 ++- .../shared/utils/url_security.py | 380 ++++++++---------- 14 files changed, 404 insertions(+), 334 deletions(-) delete mode 100644 packages/shared-python/shared/utils/outbound_url_validator.py diff --git a/apps/api/app/api/v1/routes/jobs.py b/apps/api/app/api/v1/routes/jobs.py index 2cdf7480a..31bc5268b 100644 --- a/apps/api/app/api/v1/routes/jobs.py +++ b/apps/api/app/api/v1/routes/jobs.py @@ -53,8 +53,8 @@ StandardErrorObject, ) from shared.services.storage.file_upload_service import FileUploadService -from shared.utils.outbound_url_validator import ( - validate_outbound_url_async, +from shared.utils.url_security import ( + validate_http_url_and_resolve_ip_async, ) from shared.utils.error_details import normalize_error_details from shared.utils.url_file_type import resolve_file_extension_async @@ -322,8 +322,8 @@ async def create_job( # pyright: ignore[reportGeneralTypeIssues] if payload.webhook: # Check for URL validity if payload.webhook.url: - validation_result = await validate_outbound_url_async( - payload.webhook.url + validation_result = await validate_http_url_and_resolve_ip_async( + payload.webhook.url, ) if not validation_result.is_valid: raise WebhookConfigException( diff --git a/apps/api/app/api/v1/routes/s3_events.py b/apps/api/app/api/v1/routes/s3_events.py index 40d0bb99e..ee92dec6a 100644 --- a/apps/api/app/api/v1/routes/s3_events.py +++ b/apps/api/app/api/v1/routes/s3_events.py @@ -21,7 +21,9 @@ from shared.utils.pinned_outbound_http import ( send_pinned_outbound_request, ) -from shared.utils.outbound_url_validator import validate_outbound_url_async +from shared.utils.url_security import ( + validate_http_url_and_resolve_ip_async, +) router = APIRouter(tags=["Internal"]) @@ -266,7 +268,9 @@ async def handle_sns_event(body: bytes): async def confirm_sns_subscription(subscribe_url: str) -> dict[str, str]: """Confirm an SNS subscription after SSRF validation and IP pinning.""" - validation = await validate_outbound_url_async(subscribe_url) + validation = await validate_http_url_and_resolve_ip_async( + subscribe_url, + ) if not validation.is_valid: logger.warning( f"SNS subscription confirmation URL failed validation: {validation.error_message}" diff --git a/apps/api/tests/contract/test_job_creation_contract.py b/apps/api/tests/contract/test_job_creation_contract.py index 69f9ddce4..7b4025b44 100644 --- a/apps/api/tests/contract/test_job_creation_contract.py +++ b/apps/api/tests/contract/test_job_creation_contract.py @@ -43,6 +43,7 @@ async def _load_job_record(job_id: str) -> dict[str, object]: status, source_type, s3_key, + webhook_url, webhook_enabled, job_metadata FROM jobs @@ -566,6 +567,11 @@ def apply_async( ) class _FakeCeleryApp: + def __init__(self) -> None: + from types import SimpleNamespace + + self.conf = SimpleNamespace(task_routes={}) + def signature(self, task_name: str) -> _FakeCeleryTask: return _FakeCeleryTask(task_name) @@ -653,6 +659,137 @@ def signature(self, task_name: str) -> _FakeCeleryTask: ] +@pytest.mark.asyncio +async def test_should_accept_an_http_webhook_url_when_creating_a_file_job_in_production( + monkeypatch: MonkeyPatch, + developer_api_client_factory: Callable[ + [], AbstractAsyncContextManager[AsyncClient] + ], +) -> None: + webhook_url = "http://hooks.example.test/notify" + payload: dict[str, object] = { + "namespace": "contract-jobs", + "source_type": "file", + "file_name": "contract-upload.pdf", + "data_id": "contract-job-http-webhook", + "webhook": {"url": webhook_url}, + } + + def resolve_public_address( + host: str, + port: int | None, + *args: object, + **kwargs: object, + ) -> list[tuple[socket.AddressFamily, socket.SocketKind, int, str, tuple[str, int]]]: + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 0))] + + monkeypatch.setattr(socket, "getaddrinfo", resolve_public_address) + + async with developer_api_client_factory() as api_client: + response = await api_client.post("/api/v1/jobs", json=payload) + + assert response.status_code == 200 + assert response.headers["x-request-id"] + + response_json: dict[str, object] = response.json() + job_id = cast(str, response_json["job_id"]) + + assert response_json["status"] == "waiting-file" + assert response_json["source_type"] == "file" + + job_row = await _load_job_record(job_id) + assert job_row["webhook_enabled"] is True + assert job_row["webhook_url"] == webhook_url + + +@pytest.mark.asyncio +async def test_should_accept_a_private_url_source_when_creating_a_url_job_in_local_development( + monkeypatch: MonkeyPatch, + developer_api_client_factory: Callable[ + [], AbstractAsyncContextManager[AsyncClient] + ], +) -> None: + source_url = "http://127.0.0.1/contracts/local-private.pdf" + payload: dict[str, str] = { + "namespace": "contract-jobs", + "source_type": "url", + "source_url": source_url, + "data_id": "contract-job-url-local-private-host", + } + scheduled_tasks: list[dict[str, object]] = [] + + class _FakeCeleryTask: + def __init__(self, task_name: str) -> None: + self._task_name = task_name + + def apply_async( + self, + *, + args: list[object], + kwargs: dict[str, object], + ) -> None: + scheduled_tasks.append( + { + "task_name": self._task_name, + "args": args, + "kwargs": kwargs, + } + ) + + class _FakeCeleryApp: + def __init__(self) -> None: + from types import SimpleNamespace + + self.conf = SimpleNamespace(task_routes={}) + + def signature(self, task_name: str) -> _FakeCeleryTask: + return _FakeCeleryTask(task_name) + + def resolve_private_address( + host: str, + port: int | None, + *args: object, + **kwargs: object, + ) -> list[tuple[socket.AddressFamily, socket.SocketKind, int, str, tuple[str, int]]]: + return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("127.0.0.1", 0))] + + import shared.core.celery_app as celery_app_module + monkeypatch.setattr(socket, "getaddrinfo", resolve_private_address) + monkeypatch.setattr( + celery_app_module, + "get_celery_app", + lambda: _FakeCeleryApp(), + ) + + async with developer_api_client_factory() as api_client: + import shared.core.config as shared_config_module + + monkeypatch.setattr(shared_config_module.app_config, "ENVIRONMENT", "local") + response = await api_client.post("/api/v1/jobs", json=payload) + + assert response.status_code == 200 + assert response.headers["x-request-id"] + + response_json: dict[str, object] = response.json() + job_id = cast(str, response_json["job_id"]) + + assert response_json["status"] == "waiting-file" + assert response_json["source_type"] == "url" + + job_row = await _load_job_record(job_id) + job_metadata = cast(dict[str, object], job_row["job_metadata"]) + + assert job_row["source_type"] == "url" + assert job_metadata["source_url"] == source_url + assert scheduled_tasks == [ + { + "task_name": "app.core.tasks.kb_tasks.upload_url_file_task", + "args": [job_id, source_url, "local-dev-user"], + "kwargs": {"job_type": "kb_management"}, + } + ] + + @pytest.mark.asyncio async def test_should_reject_url_source_when_url_resolves_to_private_network( monkeypatch: MonkeyPatch, @@ -733,7 +870,6 @@ async def head( return _FakeHeadResponse() import shared.utils.http_clients as http_clients_module - monkeypatch.setattr( http_clients_module, "get_async_client", diff --git a/apps/worker/app/services/storage/sync_storage_service.py b/apps/worker/app/services/storage/sync_storage_service.py index d378c9451..268c68b90 100644 --- a/apps/worker/app/services/storage/sync_storage_service.py +++ b/apps/worker/app/services/storage/sync_storage_service.py @@ -14,7 +14,7 @@ from shared.core.config.storage import get_cached_storage_adapter from shared.core.exceptions.domain_exceptions import StorageServiceException from shared.utils.pinned_outbound_http import download_pinned_outbound_file -from shared.utils.url_security import validate_public_http_url_and_resolve_ip +from shared.utils.url_security import validate_http_url_and_resolve_ip def get_storage_adapter(): @@ -111,10 +111,13 @@ def download_file_from_url(file_url: str) -> str: """Download a URL file through SSRF validation and IP pinning.""" temp_file_path = "" try: - validation = validate_public_http_url_and_resolve_ip( - file_url, - field="source_url", - ) + validation = validate_http_url_and_resolve_ip(file_url) + if not validation.is_valid or not validation.validated_ip: + raise StorageServiceException( + internal_message=f"Invalid URL: {validation.error_message}", + operation="download_from_url", + ) + temp_dir = getattr(settings, "TMP_PATH", "/tmp") os.makedirs(temp_dir, exist_ok=True) download_result = download_pinned_outbound_file( @@ -126,6 +129,8 @@ def download_file_from_url(file_url: str) -> str: ) temp_file_path = download_result.temp_file_path return temp_file_path + except StorageServiceException: + raise except Exception as e: if os.path.exists(temp_file_path): os.remove(temp_file_path) diff --git a/apps/worker/tests/contract/test_url_upload_contract.py b/apps/worker/tests/contract/test_url_upload_contract.py index 6a4674403..8c3ade508 100644 --- a/apps/worker/tests/contract/test_url_upload_contract.py +++ b/apps/worker/tests/contract/test_url_upload_contract.py @@ -3,7 +3,6 @@ import os import socket from pathlib import Path -from types import SimpleNamespace from typing import Any from uuid import uuid4 diff --git a/apps/worker/tests/contract/test_webhook_recovery_contract.py b/apps/worker/tests/contract/test_webhook_recovery_contract.py index 3afab298f..dc9266886 100644 --- a/apps/worker/tests/contract/test_webhook_recovery_contract.py +++ b/apps/worker/tests/contract/test_webhook_recovery_contract.py @@ -119,8 +119,8 @@ def publish(self, **kwargs: Any) -> SimpleNamespace: ) monkeypatch.setattr( qstash_publisher, - "validate_outbound_url", - lambda url: SimpleNamespace( + "validate_http_url_and_resolve_ip", + lambda *args, **kwargs: SimpleNamespace( is_valid=True, error_message=None, validated_ip="93.184.216.34", diff --git a/packages/shared-python/shared/models/database/job.py b/packages/shared-python/shared/models/database/job.py index 2613b8722..76c9f1df1 100644 --- a/packages/shared-python/shared/models/database/job.py +++ b/packages/shared-python/shared/models/database/job.py @@ -29,10 +29,10 @@ if TYPE_CHECKING: + from shared.models.database.job_result import JobResult + from shared.models.database.job_state_audit_log import JobStateAuditLog from shared.models.database.job_state_history import JobStateHistory from shared.models.database.webhook_log import WebhookLog - from shared.models.database.job_state_audit_log import JobStateAuditLog - from shared.models.database.job_result import JobResult class Job(Base): diff --git a/packages/shared-python/shared/services/storage/file_upload_service.py b/packages/shared-python/shared/services/storage/file_upload_service.py index e6398a493..c60f3c557 100644 --- a/packages/shared-python/shared/services/storage/file_upload_service.py +++ b/packages/shared-python/shared/services/storage/file_upload_service.py @@ -15,9 +15,7 @@ from shared.utils.pinned_outbound_http import ( download_pinned_outbound_file_async, ) -from shared.utils.url_security import ( - validate_public_http_url_and_resolve_ip_async, -) +from shared.utils.url_security import validate_http_url_and_resolve_ip_async class FileUploadService: @@ -432,10 +430,13 @@ async def _download_file_from_url(self, file_url: str) -> str: """Download a file from a URL into a temporary directory.""" temp_file_path = "" try: - validation = await validate_public_http_url_and_resolve_ip_async( - file_url, - field="source_url", - ) + validation = await validate_http_url_and_resolve_ip_async(file_url) + if not validation.is_valid or not validation.validated_ip: + raise StorageServiceException( + internal_message=f"Invalid URL: {validation.error_message}", + operation="download_from_url", + ) + temp_dir = getattr(settings, "TMP_PATH", "/tmp") os.makedirs(temp_dir, exist_ok=True) download_result = await download_pinned_outbound_file_async( diff --git a/packages/shared-python/shared/services/webhook/dispatcher.py b/packages/shared-python/shared/services/webhook/dispatcher.py index c3d166e58..e5ee4610d 100644 --- a/packages/shared-python/shared/services/webhook/dispatcher.py +++ b/packages/shared-python/shared/services/webhook/dispatcher.py @@ -32,9 +32,9 @@ from shared.utils.pinned_outbound_http import ( send_pinned_outbound_request, ) -from shared.utils.outbound_url_validator import ( - OutboundURLValidationResult, - validate_outbound_url_async, +from shared.utils.url_security import ( + HTTPURLValidationResult, + validate_http_url_and_resolve_ip_async, ) # Constants @@ -166,8 +166,8 @@ async def _send_webhook( attempt_id = str(uuid.uuid4()) # SSRF Protection - validation: OutboundURLValidationResult = await validate_outbound_url_async( - event.target_url + validation: HTTPURLValidationResult = await validate_http_url_and_resolve_ip_async( + event.target_url, ) if not validation.is_valid: logger.warning( diff --git a/packages/shared-python/shared/services/webhook/qstash_publisher.py b/packages/shared-python/shared/services/webhook/qstash_publisher.py index 065601a7a..a3c52659c 100644 --- a/packages/shared-python/shared/services/webhook/qstash_publisher.py +++ b/packages/shared-python/shared/services/webhook/qstash_publisher.py @@ -22,7 +22,9 @@ from shared.core.config import app_config from shared.core.exceptions.domain_exceptions import QStashServiceException from shared.models.database.webhook import WebhookEventStatus -from shared.utils.outbound_url_validator import validate_outbound_url +from shared.utils.url_security import ( + validate_http_url_and_resolve_ip, +) class QStashWebhookPublisher: @@ -84,7 +86,9 @@ def publish_event(self, event_id: str) -> Optional[str]: return None # SSRF pre-validation - validation = validate_outbound_url(event.target_url) + validation = validate_http_url_and_resolve_ip( + event.target_url, + ) if not validation.is_valid: logger.warning( f"QStash publish: SSRF validation failed for event {event_id}: " diff --git a/packages/shared-python/shared/testing/contract_runtime.py b/packages/shared-python/shared/testing/contract_runtime.py index 5453c3784..11aa49efb 100644 --- a/packages/shared-python/shared/testing/contract_runtime.py +++ b/packages/shared-python/shared/testing/contract_runtime.py @@ -302,7 +302,7 @@ def configure_contract_environment( _reset_contract_storage_state(database_url) environment: dict[str, str] = { - "ENVIRONMENT": "development", + "ENVIRONMENT": "production", "API_STANDALONE_MODE_ENABLED": "true", "WEBHOOK_MASTER_KEY": CONTRACT_WEBHOOK_MASTER_KEY, "DATABASE_URL": database_url, diff --git a/packages/shared-python/shared/utils/outbound_url_validator.py b/packages/shared-python/shared/utils/outbound_url_validator.py deleted file mode 100644 index 40a203b86..000000000 --- a/packages/shared-python/shared/utils/outbound_url_validator.py +++ /dev/null @@ -1,91 +0,0 @@ -""" -Outbound URL Validation Utilities - -Shared SSRF protection for outbound HTTP targets via DNS/IP validation + IP pinning. -""" - -from dataclasses import dataclass -from typing import Optional -from urllib.parse import urlparse - -from shared.core.config import app_config -from shared.utils.url_security import resolve_public_hostname, resolve_public_hostname_async - - -@dataclass -class OutboundURLValidationResult: - """Result of outbound URL validation, including a pinned IP address.""" - - is_valid: bool - error_message: Optional[str] = None - validated_ip: Optional[str] = None - hostname: Optional[str] = None - - -async def validate_outbound_url_async(url: str) -> OutboundURLValidationResult: - """ - Async outbound URL validation with SSRF protection and IP pinning. - - Returns OutboundURLValidationResult with a pinned IP address, eliminating - the DNS rebinding TOCTOU window for later outbound requests. - """ - try: - parsed = urlparse(url) - is_dev: bool = app_config.ENVIRONMENT.lower() in ("dev", "development", "local") - allowed_schemes: list[str] = ["https"] if not is_dev else ["https", "http"] - if parsed.scheme not in allowed_schemes: - return OutboundURLValidationResult( - is_valid=False, - error_message=f"Invalid scheme: {parsed.scheme}. Must be HTTPS.", - ) - hostname: Optional[str] = parsed.hostname - if not hostname: - return OutboundURLValidationResult( - is_valid=False, error_message="URL must have a hostname" - ) - validated_ip: str = await resolve_public_hostname_async(hostname) - return OutboundURLValidationResult( - is_valid=True, - validated_ip=validated_ip, - hostname=hostname, - ) - except ValueError as exc: - return OutboundURLValidationResult(is_valid=False, error_message=str(exc)) - except Exception as exc: - return OutboundURLValidationResult( - is_valid=False, - error_message=f"URL validation failed: {exc}", - ) - - -def validate_outbound_url(url: str) -> OutboundURLValidationResult: - """Sync outbound URL validation with SSRF checks.""" - try: - parsed = urlparse(url) - is_dev: bool = app_config.ENVIRONMENT.lower() in ("dev", "development", "local") - allowed_schemes: list[str] = ["https"] if not is_dev else ["https", "http"] - if parsed.scheme not in allowed_schemes: - return OutboundURLValidationResult( - is_valid=False, - error_message=f"Invalid scheme: {parsed.scheme}. Must be HTTPS.", - ) - - hostname: Optional[str] = parsed.hostname - if not hostname: - return OutboundURLValidationResult( - is_valid=False, error_message="URL must have a hostname" - ) - - validated_ip: str = resolve_public_hostname(hostname) - return OutboundURLValidationResult( - is_valid=True, - validated_ip=validated_ip, - hostname=hostname, - ) - except ValueError as exc: - return OutboundURLValidationResult(is_valid=False, error_message=str(exc)) - except Exception as exc: - return OutboundURLValidationResult( - is_valid=False, - error_message=f"URL validation failed: {exc}", - ) diff --git a/packages/shared-python/shared/utils/url_file_type.py b/packages/shared-python/shared/utils/url_file_type.py index 0a0793a18..7314d7e23 100644 --- a/packages/shared-python/shared/utils/url_file_type.py +++ b/packages/shared-python/shared/utils/url_file_type.py @@ -6,17 +6,16 @@ """ import os -from urllib.parse import urlparse +from urllib.parse import urljoin, urlparse from loguru import logger from shared.core.config import settings from shared.core.exceptions.domain_exceptions import ValidationException from shared.utils.url_security import ( - MAX_SAFE_REDIRECTS, + HTTPURLValidationResult, SafePublicHTTPURL, - get_safe_public_http_url, - validate_public_http_redirect_url, + validate_http_url_and_resolve_ip, ) # Content-Type to file extension mapping @@ -40,6 +39,45 @@ } REDIRECT_STATUS_CODES: set[int] = {301, 302, 303, 307, 308} +MAX_SAFE_REDIRECTS: int = 5 +URL_VALIDATION_DESCRIPTIONS: dict[str, str] = { + "unsupported_scheme": "URL must use http or https", + "missing_hostname": "URL must include a hostname", + "hostname_resolution_failed": "URL hostname could not be resolved", + "invalid_resolved_address": "URL resolved to an invalid IP", + "hostname_not_allowed": "URL host is not allowed", +} + + +def _validate_source_url(url: str, field: str) -> SafePublicHTTPURL: + validation = validate_http_url_and_resolve_ip(url) + if not validation.is_valid: + _raise_url_validation_error(field, validation) + return SafePublicHTTPURL(url) + + +def _validate_redirect_url( + url: str, + redirect_url: str, + field: str, +) -> SafePublicHTTPURL: + return _validate_source_url(urljoin(url, redirect_url), field=field) + + +def _raise_url_validation_error( + field: str, + validation: HTTPURLValidationResult, +) -> None: + description = URL_VALIDATION_DESCRIPTIONS.get( + str(validation.failure_reason), + "URL is invalid", + ) + + raise ValidationException( + user_message="Invalid URL", + violations=[{"field": field, "description": description}], + internal_message=validation.error_message, + ) def _extension_from_path(url: str) -> str | None: @@ -71,7 +109,7 @@ async def resolve_file_extension_async(url: str) -> str | None: 2. If that fails, send a HEAD request and read Content-Type. 3. Return None if neither method produces a supported extension. """ - safe_url = get_safe_public_http_url(url, field="source_url") + safe_url = _validate_source_url(url, field="source_url") ext = _extension_from_path(safe_url) if ext: @@ -91,7 +129,7 @@ async def resolve_file_extension_async(url: str) -> str | None: location = response.headers.get("location") if not location: break - request_url = validate_public_http_redirect_url( + request_url = _validate_redirect_url( request_url, location, field="source_url", @@ -122,7 +160,7 @@ def resolve_file_extension_sync(url: str) -> str | None: Same logic as async variant but uses the shared sync httpx client. """ - safe_url = get_safe_public_http_url(url, field="source_url") + safe_url = _validate_source_url(url, field="source_url") ext = _extension_from_path(safe_url) if ext: @@ -142,7 +180,7 @@ def resolve_file_extension_sync(url: str) -> str | None: location = response.headers.get("location") if not location: break - request_url = validate_public_http_redirect_url( + request_url = _validate_redirect_url( request_url, location, field="source_url", diff --git a/packages/shared-python/shared/utils/url_security.py b/packages/shared-python/shared/utils/url_security.py index 0b76c06be..bc54c1107 100644 --- a/packages/shared-python/shared/utils/url_security.py +++ b/packages/shared-python/shared/utils/url_security.py @@ -2,38 +2,25 @@ import ipaddress import socket from dataclasses import dataclass -from typing import cast -from urllib.parse import urljoin, urlparse - -from shared.core.exceptions.domain_exceptions import ValidationException - -AddressInfo = tuple[int, int, int, str, tuple[str, ...]] -AllowedIPAddress = ipaddress.IPv4Address | ipaddress.IPv6Address -MAX_SAFE_REDIRECTS = 5 - - -class URLSecurityError(ValueError): - """Base error for URL safety validation failures.""" - - -class HostnameResolutionError(URLSecurityError): - """Raised when a hostname cannot be resolved.""" - - -class HostnameNotAllowedError(URLSecurityError): - """Raised when a hostname resolves to a blocked network address.""" - - -class InvalidResolvedAddressError(URLSecurityError): - """Raised when DNS returns an invalid IP address.""" - - -class UnsupportedURLSchemeError(URLSecurityError): - """Raised when a URL uses a disallowed scheme.""" - - -class MissingURLHostnameError(URLSecurityError): - """Raised when a URL does not include a hostname.""" +from typing import Literal, TypeAlias, cast +from urllib.parse import urlparse + +SocketAddress: TypeAlias = tuple[object, ...] +AddressInfo: TypeAlias = tuple[int, int, int, str, SocketAddress] +AllowedIPAddress: TypeAlias = ipaddress.IPv4Address | ipaddress.IPv6Address +URLValidationFailureReason: TypeAlias = Literal[ + "unsupported_scheme", + "missing_hostname", + "hostname_resolution_failed", + "invalid_resolved_address", + "hostname_not_allowed", + "validation_failed", +] +HTTP_SCHEMES: frozenset[str] = frozenset({"http", "https"}) +DEVELOPMENT_ENVIRONMENTS: frozenset[str] = frozenset({"dev", "development", "local"}) +BLOCKED_HOSTNAMES: frozenset[str] = frozenset( + {"localhost", "ip6-localhost", "ip6-loopback"} +) class SafePublicHTTPURL(str): @@ -41,186 +28,107 @@ class SafePublicHTTPURL(str): @dataclass(frozen=True) -class PublicHTTPURLValidationResult: - """Validated public HTTP URL and its pinned public IP address.""" - - url: SafePublicHTTPURL - validated_ip: str - - -def validate_public_http_url(url: str, field: str = "url") -> None: - """Reject URL inputs that could target internal networks or local services.""" - try: - _validate_public_http_url(url) - except UnsupportedURLSchemeError as exc: - raise _build_url_validation_error(field, "URL must use http or https") from exc - except MissingURLHostnameError as exc: - raise _build_url_validation_error(field, "URL must include a hostname") from exc - except HostnameResolutionError as exc: - raise _build_url_validation_error(field, "URL hostname could not be resolved") from exc - except InvalidResolvedAddressError as exc: - raise _build_url_validation_error(field, "URL resolved to an invalid IP") from exc - except HostnameNotAllowedError as exc: - raise _build_url_validation_error(field, "URL host is not allowed") from exc - - -def validate_public_http_redirect_url( - url: str, - redirect_url: str, - field: str = "url", -) -> SafePublicHTTPURL: - """Resolve and validate an HTTP redirect target before following it.""" - resolved_url = urljoin(url, redirect_url) - validate_public_http_url(resolved_url, field=field) - return SafePublicHTTPURL(resolved_url) - +class HTTPURLValidationResult: + """Validated HTTP URL and its pinned IP address.""" -def get_safe_public_http_url(url: str, field: str = "url") -> SafePublicHTTPURL: - """Return a validated public HTTP URL for outbound HTTP clients.""" - validate_public_http_url(url, field=field) - return SafePublicHTTPURL(url) + is_valid: bool + url: str + hostname: str | None = None + validated_ip: str | None = None + error_message: str | None = None + failure_reason: URLValidationFailureReason | None = None -def validate_public_http_url_and_resolve_ip( +def validate_http_url_and_resolve_ip( url: str, - field: str = "url", -) -> PublicHTTPURLValidationResult: - """Validate a public HTTP URL and return the IP selected during validation.""" +) -> HTTPURLValidationResult: + """Validate an HTTP URL policy and return the pinned resolved IP address.""" try: - validated_ip = _validate_public_http_url(url) - return PublicHTTPURLValidationResult( - url=SafePublicHTTPURL(url), - validated_ip=validated_ip, - ) - except UnsupportedURLSchemeError as exc: - raise _build_url_validation_error(field, "URL must use http or https") from exc - except MissingURLHostnameError as exc: - raise _build_url_validation_error(field, "URL must include a hostname") from exc - except HostnameResolutionError as exc: - raise _build_url_validation_error(field, "URL hostname could not be resolved") from exc - except InvalidResolvedAddressError as exc: - raise _build_url_validation_error(field, "URL resolved to an invalid IP") from exc - except HostnameNotAllowedError as exc: - raise _build_url_validation_error(field, "URL host is not allowed") from exc - - -async def validate_public_http_url_and_resolve_ip_async( - url: str, - field: str = "url", -) -> PublicHTTPURLValidationResult: - """Validate a public HTTP URL asynchronously and return the selected IP.""" - try: - parsed_url = urlparse(url) - if parsed_url.scheme not in {"http", "https"}: - raise UnsupportedURLSchemeError( - f"Unsupported URL scheme: {parsed_url.scheme}" + can_use_private_hosts = _can_use_private_hosts() + hostname_result = _validate_url_hostname(url) + if not hostname_result.is_valid: + return hostname_result + + hostname = cast(str, hostname_result.hostname) + if not can_use_private_hosts and _is_blocked_hostname(hostname): + return _build_url_validation_failure( + url, + f"Hostname {hostname} failed validation", + "hostname_not_allowed", ) - hostname = parsed_url.hostname - if not hostname: - raise MissingURLHostnameError("URL must include a hostname") + try: + address_infos = cast(list[AddressInfo], socket.getaddrinfo(hostname, None)) + except socket.gaierror as exc: + return _build_url_validation_failure( + url, + f"Unable to resolve hostname {hostname}: {exc}", + "hostname_resolution_failed", + ) - validated_ip = await resolve_public_hostname_async(hostname) - return PublicHTTPURLValidationResult( - url=SafePublicHTTPURL(url), - validated_ip=validated_ip, + return _validate_resolved_addresses( + url, + hostname, + address_infos, + can_use_private_hosts=can_use_private_hosts, ) - except UnsupportedURLSchemeError as exc: - raise _build_url_validation_error(field, "URL must use http or https") from exc - except MissingURLHostnameError as exc: - raise _build_url_validation_error(field, "URL must include a hostname") from exc - except HostnameResolutionError as exc: - raise _build_url_validation_error(field, "URL hostname could not be resolved") from exc - except InvalidResolvedAddressError as exc: - raise _build_url_validation_error(field, "URL resolved to an invalid IP") from exc - except HostnameNotAllowedError as exc: - raise _build_url_validation_error(field, "URL host is not allowed") from exc - - -def get_safe_redirect_url(url: str, redirect_url: str) -> str: - """Resolve and validate an HTTP redirect target for internal network callers.""" - resolved_url = urljoin(url, redirect_url) - _validate_public_http_url(resolved_url) - return resolved_url - - -def resolve_public_hostname(hostname: str) -> str: - """Resolve a hostname to a public IP address and return the pinned IP.""" - _ensure_hostname_is_allowed(hostname) - try: - address_infos = cast(list[AddressInfo], socket.getaddrinfo(hostname, None)) - except socket.gaierror as exc: - raise HostnameResolutionError( - f"Unable to resolve hostname {hostname}: {exc}" - ) from exc - - return _select_public_ip_address(hostname, address_infos) - - -def _validate_public_http_url(url: str) -> str: - parsed_url = urlparse(url) - if parsed_url.scheme not in {"http", "https"}: - raise UnsupportedURLSchemeError(f"Unsupported URL scheme: {parsed_url.scheme}") - - hostname = parsed_url.hostname - if not hostname: - raise MissingURLHostnameError("URL must include a hostname") - - return resolve_public_hostname(hostname) - - -async def resolve_public_hostname_async(hostname: str) -> str: - """Resolve a hostname to a public IP address asynchronously.""" - _ensure_hostname_is_allowed(hostname) - try: - loop = asyncio.get_running_loop() - address_infos = cast( - list[AddressInfo], - await loop.getaddrinfo(hostname, None), + except Exception as exc: + return _build_url_validation_failure( + url, + f"URL validation failed: {exc}", + "validation_failed", ) - except socket.gaierror as exc: - raise HostnameResolutionError( - f"Unable to resolve hostname {hostname}: {exc}" - ) from exc - return _select_public_ip_address(hostname, address_infos) - -def is_public_ip_address(address: str) -> bool: - """Return whether an IP string is safe to use as a public network target.""" +async def validate_http_url_and_resolve_ip_async( + url: str, +) -> HTTPURLValidationResult: + """Validate an HTTP URL policy asynchronously and return the pinned IP address.""" try: - ip_address = ipaddress.ip_address(address) - except ValueError: - return False - return _is_public_ip_address(ip_address) - + can_use_private_hosts = _can_use_private_hosts() + hostname_result = _validate_url_hostname(url) + if not hostname_result.is_valid: + return hostname_result + + hostname = cast(str, hostname_result.hostname) + if not can_use_private_hosts and _is_blocked_hostname(hostname): + return _build_url_validation_failure( + url, + f"Hostname {hostname} failed validation", + "hostname_not_allowed", + ) -def _ensure_hostname_is_allowed(hostname: str) -> None: - if _is_blocked_hostname(hostname): - raise HostnameNotAllowedError(f"Hostname {hostname} failed validation") + try: + loop = asyncio.get_running_loop() + address_infos = cast( + list[AddressInfo], + await loop.getaddrinfo(hostname, None), + ) + except socket.gaierror as exc: + return _build_url_validation_failure( + url, + f"Unable to resolve hostname {hostname}: {exc}", + "hostname_resolution_failed", + ) + return _validate_resolved_addresses( + url, + hostname, + address_infos, + can_use_private_hosts=can_use_private_hosts, + ) + except Exception as exc: + return _build_url_validation_failure( + url, + f"URL validation failed: {exc}", + "validation_failed", + ) -def _select_public_ip_address(hostname: str, address_infos: list[AddressInfo]) -> str: - resolved_addresses = _extract_resolved_addresses(address_infos) - if not resolved_addresses: - raise HostnameResolutionError(f"Unable to resolve hostname {hostname}") - selected_address: str | None = None - for address in resolved_addresses: - try: - ip_address = ipaddress.ip_address(address) - except ValueError as exc: - raise InvalidResolvedAddressError( - f"Hostname {hostname} resolved to invalid IP {address}" - ) from exc - if not _is_public_ip_address(ip_address): - raise HostnameNotAllowedError(f"Hostname {hostname} failed validation") - if selected_address is None: - selected_address = address +def _can_use_private_hosts() -> bool: + from shared.core.config import app_config - if selected_address: - return selected_address - raise HostnameNotAllowedError(f"Hostname {hostname} failed validation") + return app_config.ENVIRONMENT.lower() in DEVELOPMENT_ENVIRONMENTS def _extract_resolved_addresses(address_infos: list[AddressInfo]) -> list[str]: @@ -232,6 +140,8 @@ def _extract_resolved_addresses(address_infos: list[AddressInfo]) -> list[str]: if not socket_address: continue address = socket_address[0] + if not isinstance(address, str): + continue if address and address not in seen_addresses: resolved_addresses.append(address) seen_addresses.add(address) @@ -240,11 +150,9 @@ def _extract_resolved_addresses(address_infos: list[AddressInfo]) -> list[str]: def _is_blocked_hostname(hostname: str) -> bool: normalized_hostname = hostname.rstrip(".").lower() - if normalized_hostname in {"localhost", "ip6-localhost", "ip6-loopback"}: - return True - if normalized_hostname.endswith(".localhost") or normalized_hostname.endswith(".local"): - return True - return False + return normalized_hostname in BLOCKED_HOSTNAMES or normalized_hostname.endswith( + (".localhost", ".local") + ) def _is_public_ip_address(ip_address: AllowedIPAddress) -> bool: @@ -262,8 +170,74 @@ def _is_blocked_ip_address(ip_address: AllowedIPAddress) -> bool: ) -def _build_url_validation_error(field: str, description: str) -> ValidationException: - return ValidationException( - user_message="Invalid URL", - violations=[{"field": field, "description": description}], +def _validate_url_hostname(url: str) -> HTTPURLValidationResult: + parsed_url = urlparse(url) + if parsed_url.scheme not in HTTP_SCHEMES: + return _build_url_validation_failure( + url, + f"Unsupported URL scheme: {parsed_url.scheme}", + "unsupported_scheme", + ) + + hostname = parsed_url.hostname + if not hostname: + return _build_url_validation_failure( + url, + "URL must include a hostname", + "missing_hostname", + ) + + return HTTPURLValidationResult(is_valid=True, url=url, hostname=hostname) + + +def _validate_resolved_addresses( + url: str, + hostname: str, + address_infos: list[AddressInfo], + *, + can_use_private_hosts: bool, +) -> HTTPURLValidationResult: + resolved_addresses = _extract_resolved_addresses(address_infos) + if not resolved_addresses: + return _build_url_validation_failure( + url, + f"Unable to resolve hostname {hostname}", + "hostname_resolution_failed", + ) + + for address in resolved_addresses: + try: + ip_address = ipaddress.ip_address(address) + except ValueError: + return _build_url_validation_failure( + url, + f"Hostname {hostname} resolved to invalid IP {address}", + "invalid_resolved_address", + ) + + if not can_use_private_hosts and not _is_public_ip_address(ip_address): + return _build_url_validation_failure( + url, + f"Hostname {hostname} failed validation", + "hostname_not_allowed", + ) + + return HTTPURLValidationResult( + is_valid=True, + url=url, + hostname=hostname, + validated_ip=resolved_addresses[0], + ) + + +def _build_url_validation_failure( + url: str, + error_message: str, + failure_reason: URLValidationFailureReason, +) -> HTTPURLValidationResult: + return HTTPURLValidationResult( + is_valid=False, + url=url, + error_message=error_message, + failure_reason=failure_reason, ) From 6b2079c9d8bcc67d1ac32000edccf9e85a0abe14 Mon Sep 17 00:00:00 2001 From: suguanYang Date: Thu, 7 May 2026 11:27:47 +0800 Subject: [PATCH 32/32] fix: update environment assertion in version payload test to production --- apps/api/tests/contract/test_version_contract.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/apps/api/tests/contract/test_version_contract.py b/apps/api/tests/contract/test_version_contract.py index 14090dbbb..7b8cc599a 100644 --- a/apps/api/tests/contract/test_version_contract.py +++ b/apps/api/tests/contract/test_version_contract.py @@ -21,7 +21,7 @@ async def test_should_return_version_payload_for_the_v1_version_endpoint( assert response_json["version"] assert "commit" in response_json assert "build_time" in response_json - assert response_json["environment"] == "development" + assert response_json["environment"] == "production" @pytest.mark.asyncio