From 457b273c36a38ff5b09e26d9e2f79523bd7e7913 Mon Sep 17 00:00:00 2001 From: Ayush8923 <80516839+Ayush8923@users.noreply.github.com> Date: Sat, 8 Aug 2026 15:43:52 +0530 Subject: [PATCH] feat(onboarding): type safety checks --- backend/app/api/routes/credentials.py | 50 +++--- backend/app/api/routes/onboarding.py | 2 +- backend/app/api/routes/project.py | 15 +- backend/app/api/routes/users.py | 19 +-- backend/app/core/providers.py | 185 ++++++++++++++++++--- backend/app/crud/credentials.py | 30 +--- backend/app/crud/onboarding.py | 16 +- backend/app/models/credentials.py | 72 +++++++- backend/app/models/onboarding.py | 57 +++---- backend/app/models/project.py | 3 +- backend/app/tests/api/routes/test_creds.py | 7 +- backend/app/tests/crud/test_credentials.py | 49 ++---- docs/wiki/modules/platform.md | 23 +++ docs/wiki/modules/tenancy.md | 9 + 14 files changed, 362 insertions(+), 175 deletions(-) diff --git a/backend/app/api/routes/credentials.py b/backend/app/api/routes/credentials.py index 8ee344546..dd8140902 100644 --- a/backend/app/api/routes/credentials.py +++ b/backend/app/api/routes/credentials.py @@ -5,7 +5,11 @@ from app.api.deps import AuthContextDep, SessionDep from app.api.permissions import Permission, require_permission from app.core.exception_handlers import HTTPException -from app.core.providers import mask_credential_fields, validate_provider +from app.core.providers import ( + CredentialPayload, + mask_credential_fields, + validate_provider, +) from app.crud.credentials import ( get_creds_by_org, get_provider_credential, @@ -14,7 +18,7 @@ set_creds_for_org, update_creds_for_org, ) -from app.models import CredsCreate, CredsPublic, CredsUpdate +from app.models import CredsCreate, CredsPublic, CredsUpdate, Message from app.utils import APIResponse, load_description logger = logging.getLogger(__name__) @@ -32,7 +36,7 @@ def create_new_credential( session: SessionDep, creds_in: CredsCreate, _current_user: AuthContextDep, -): +) -> APIResponse[list[CredsPublic]]: # Project comes from API key context; no cross-org check needed here # Database unique constraint ensures no duplicate credentials per provider-org-project combination @@ -61,7 +65,7 @@ def read_credential( *, session: SessionDep, _current_user: AuthContextDep, -): +) -> APIResponse[list[CredsPublic]]: creds = get_creds_by_org( session=session, org_id=_current_user.organization_.id, @@ -73,7 +77,7 @@ def read_credential( @router.get( "/provider/{provider}", - response_model=APIResponse[dict | None], + response_model=APIResponse[CredentialPayload | None], description=load_description("credentials/get_provider.md"), dependencies=[Depends(require_permission(Permission.REQUIRE_PROJECT))], ) @@ -82,7 +86,7 @@ def read_provider_credential( session: SessionDep, provider: str, _current_user: AuthContextDep, -): +) -> APIResponse[CredentialPayload | None]: try: provider_enum = validate_provider(provider) except ValueError as e: @@ -113,7 +117,7 @@ def update_credential( session: SessionDep, creds_in: CredsUpdate, _current_user: AuthContextDep, -): +) -> APIResponse[list[CredsPublic]]: if not creds_in or not creds_in.provider or not creds_in.credential: logger.warning( f"[update_credential] Invalid input | organization_id: {_current_user.organization_.id}, project_id: {_current_user.project_.id}" @@ -137,7 +141,7 @@ def update_credential( @router.delete( "/provider/{provider}", - response_model=APIResponse[dict], + response_model=APIResponse[Message], description=load_description("credentials/delete_provider.md"), dependencies=[Depends(require_permission(Permission.REQUIRE_PROJECT))], ) @@ -146,7 +150,7 @@ def delete_provider_credential( session: SessionDep, provider: str, _current_user: AuthContextDep, -): +) -> APIResponse[Message]: try: provider_enum = validate_provider(provider) except ValueError as e: @@ -160,13 +164,13 @@ def delete_provider_credential( ) return APIResponse.success_response( - {"message": "Provider credentials removed successfully"} + Message(message="Provider credentials removed successfully") ) @router.delete( "", - response_model=APIResponse[dict], + response_model=APIResponse[Message], description=load_description("credentials/delete_all.md"), dependencies=[Depends(require_permission(Permission.REQUIRE_PROJECT))], ) @@ -174,7 +178,7 @@ def delete_all_credentials( *, session: SessionDep, _current_user: AuthContextDep, -): +) -> APIResponse[Message]: remove_creds_for_org( session=session, org_id=_current_user.organization_.id, @@ -182,7 +186,7 @@ def delete_all_credentials( ) return APIResponse.success_response( - {"message": "All credentials deleted successfully"} + Message(message="All credentials deleted successfully") ) @@ -198,7 +202,7 @@ def read_credentials_by_org_project( org_id: int, project_id: int, _current_user: AuthContextDep, -): +) -> APIResponse[list[CredsPublic]]: creds = get_creds_by_org( session=session, org_id=org_id, @@ -210,7 +214,7 @@ def read_credentials_by_org_project( @router.get( "/{org_id}/{project_id}/provider/{provider}", - response_model=APIResponse[dict | None], + response_model=APIResponse[CredentialPayload | None], description=load_description("credentials/get_provider_by_org_project.md"), dependencies=[Depends(require_permission(Permission.SUPERUSER))], ) @@ -221,7 +225,7 @@ def read_provider_credential_by_org_project( project_id: int, provider: str, _current_user: AuthContextDep, -): +) -> APIResponse[CredentialPayload | None]: try: provider_enum = validate_provider(provider) except ValueError as e: @@ -254,7 +258,7 @@ def update_credential_by_org_project( project_id: int, creds_in: CredsUpdate, _current_user: AuthContextDep, -): +) -> APIResponse[list[CredsPublic]]: if not creds_in or not creds_in.provider or not creds_in.credential: logger.warning( f"[update_credential_by_org_project] Invalid input | organization_id: {org_id}, project_id: {project_id}" @@ -277,7 +281,7 @@ def update_credential_by_org_project( @router.delete( "/{org_id}/{project_id}/provider/{provider}", - response_model=APIResponse[dict], + response_model=APIResponse[Message], description=load_description("credentials/delete_provider_by_org_project.md"), dependencies=[Depends(require_permission(Permission.SUPERUSER))], ) @@ -288,7 +292,7 @@ def delete_provider_credential_by_org_project( project_id: int, provider: str, _current_user: AuthContextDep, -): +) -> APIResponse[Message]: try: provider_enum = validate_provider(provider) except ValueError as e: @@ -302,13 +306,13 @@ def delete_provider_credential_by_org_project( ) return APIResponse.success_response( - {"message": "Provider credentials removed successfully"} + Message(message="Provider credentials removed successfully") ) @router.delete( "/{org_id}/{project_id}", - response_model=APIResponse[dict], + response_model=APIResponse[Message], description=load_description("credentials/delete_all_by_org_project.md"), dependencies=[Depends(require_permission(Permission.SUPERUSER))], ) @@ -318,7 +322,7 @@ def delete_all_credentials_by_org_project( org_id: int, project_id: int, _current_user: AuthContextDep, -): +) -> APIResponse[Message]: remove_creds_for_org( session=session, org_id=org_id, @@ -326,5 +330,5 @@ def delete_all_credentials_by_org_project( ) return APIResponse.success_response( - {"message": "All credentials deleted successfully"} + Message(message="All credentials deleted successfully") ) diff --git a/backend/app/api/routes/onboarding.py b/backend/app/api/routes/onboarding.py index c5f7f1231..e2d9e2103 100644 --- a/backend/app/api/routes/onboarding.py +++ b/backend/app/api/routes/onboarding.py @@ -19,7 +19,7 @@ def onboard_project_route( onboard_in: OnboardingRequest, session: SessionDep, -): +) -> APIResponse[OnboardingResponse]: response = onboard_project(session=session, onboard_in=onboard_in) metadata = None diff --git a/backend/app/api/routes/project.py b/backend/app/api/routes/project.py index db42a31fd..12c169446 100644 --- a/backend/app/api/routes/project.py +++ b/backend/app/api/routes/project.py @@ -35,7 +35,6 @@ router = APIRouter(prefix="/projects", tags=["Projects"]) -# Retrieve projects @router.get( "", dependencies=[Depends(require_permission(Permission.SUPERUSER))], @@ -54,7 +53,7 @@ def read_projects( True, description="Filter by active status. Pass false to list soft-deleted projects.", ), -): +) -> APIResponse[list[ProjectPublic]]: filters = [Project.is_active.is_(is_active)] if search and search.strip(): filters.append(Project.name.ilike(f"%{search.strip()}%")) @@ -76,7 +75,9 @@ def read_projects( response_model=APIResponse[ProjectPublic], description=load_description("projects/create.md"), ) -def create_new_project(*, session: SessionDep, project_in: ProjectCreate): +def create_new_project( + *, session: SessionDep, project_in: ProjectCreate +) -> APIResponse[ProjectPublic]: project = create_project(session=session, project_create=project_in) return APIResponse.success_response(project) @@ -111,7 +112,7 @@ def update_project_settings_route( response_model=APIResponse[ProjectPublic], description=load_description("projects/get.md"), ) -def read_project(*, session: SessionDep, project_id: int): +def read_project(*, session: SessionDep, project_id: int) -> APIResponse[ProjectPublic]: """ Retrieve a project by ID. """ @@ -128,7 +129,9 @@ def read_project(*, session: SessionDep, project_id: int): response_model=APIResponse[ProjectPublic], description=load_description("projects/update.md"), ) -def update_project(*, session: SessionDep, project_id: int, project_in: ProjectUpdate): +def update_project( + *, session: SessionDep, project_id: int, project_in: ProjectUpdate +) -> APIResponse[ProjectPublic]: project = get_project_by_id(session=session, project_id=project_id) if project is None: raise HTTPException(status_code=404, detail="Project not found") @@ -193,7 +196,7 @@ def delete_project_endpoint( session: SessionDep, project_id: int, body: DeleteRequest | None = None, -): +) -> APIResponse[None]: project = get_project_by_id(session=session, project_id=project_id) if project is None: raise HTTPException(status_code=404, detail="Project not found") diff --git a/backend/app/api/routes/users.py b/backend/app/api/routes/users.py index eff2a1a72..b13ce7856 100644 --- a/backend/app/api/routes/users.py +++ b/backend/app/api/routes/users.py @@ -1,5 +1,4 @@ import logging -from typing import Any from fastapi import APIRouter, Depends from sqlmodel import func, select @@ -37,7 +36,7 @@ response_model=UsersPublic, include_in_schema=False, ) -def read_users(session: SessionDep, skip: int = 0, limit: int = 100) -> Any: +def read_users(session: SessionDep, skip: int = 0, limit: int = 100) -> UsersPublic: count = session.exec(select(func.count()).select_from(User)).one() users = session.exec(select(User).offset(skip).limit(limit)).all() return UsersPublic(data=users, count=count) @@ -49,7 +48,7 @@ def read_users(session: SessionDep, skip: int = 0, limit: int = 100) -> Any: response_model=UserPublic, include_in_schema=False, ) -def create_user_endpoint(*, session: SessionDep, user_in: UserCreate) -> Any: +def create_user_endpoint(*, session: SessionDep, user_in: UserCreate) -> User: if get_user_by_email(session=session, email=user_in.email): logger.warning( f"[create_user_endpoint] Attempting to create user with existing email | email: {user_in.email}" @@ -76,7 +75,7 @@ def create_user_endpoint(*, session: SessionDep, user_in: UserCreate) -> Any: @router.patch("/me", response_model=UserPublic) def update_user_me( *, session: SessionDep, user_in: UserUpdateMe, current_user_dep: AuthContextDep -) -> Any: +) -> User: current_user = current_user_dep.user if user_in.email: existing_user = get_user_by_email(session=session, email=user_in.email) @@ -99,7 +98,7 @@ def update_user_me( @router.patch("/me/password", response_model=Message) def update_password_me( *, session: SessionDep, body: UpdatePassword, current_user_dep: AuthContextDep -) -> Any: +) -> Message: current_user = current_user_dep.user if not verify_password(body.current_password, current_user.hashed_password): raise HTTPException(status_code=400, detail="Incorrect password") @@ -123,7 +122,7 @@ def update_password_me( def read_user_me( session: SessionDep, current_user_dep: AuthContextDep, -) -> Any: +) -> UserPublic: user = current_user_dep.user features: list[str] = [] org = current_user_dep.organization @@ -139,7 +138,7 @@ def read_user_me( @router.delete("/me", response_model=Message) -def delete_user_me(session: SessionDep, current_user_dep: AuthContextDep) -> Any: +def delete_user_me(session: SessionDep, current_user_dep: AuthContextDep) -> Message: current_user = current_user_dep.user if current_user.is_superuser: logger.warning( @@ -159,7 +158,7 @@ def delete_user_me(session: SessionDep, current_user_dep: AuthContextDep) -> Any dependencies=[Depends(require_permission(Permission.SUPERUSER))], response_model=UserPublic, ) -def register_user(session: SessionDep, user_in: UserRegister) -> Any: +def register_user(session: SessionDep, user_in: UserRegister) -> User: """ This endpoint allows the registration of a new user and is accessible only by a superuser. """ @@ -179,7 +178,7 @@ def register_user(session: SessionDep, user_in: UserRegister) -> Any: @router.get("/{user_id}", response_model=UserPublic, include_in_schema=False) def read_user_by_id( user_id: int, session: SessionDep, current_user: AuthContextDep -) -> Any: +) -> User | None: user = session.get(User, user_id) if user == current_user.user: return user @@ -204,7 +203,7 @@ def update_user_endpoint( session: SessionDep, user_id: int, user_in: UserUpdate, -) -> Any: +) -> User: db_user = session.get(User, user_id) if not db_user: raise HTTPException( diff --git a/backend/app/core/providers.py b/backend/app/core/providers.py index 743ee1883..174fb440a 100644 --- a/backend/app/core/providers.py +++ b/backend/app/core/providers.py @@ -1,8 +1,10 @@ import logging -from typing import Any, Dict, List +from typing import Annotated, Any from enum import Enum from dataclasses import dataclass, field +from pydantic import BaseModel, ConfigDict, Field, JsonValue, ValidationError + logger = logging.getLogger(__name__) @@ -20,46 +22,145 @@ class Provider(str, Enum): PROXY = "proxy" +CredentialPayload = dict[str, JsonValue] + + +class ProviderCredentialsBase(BaseModel): + """Base for provider credential payloads. + + Extra keys are allowed and preserved: callers pass provider knobs Kaapi + does not itself require (an ``openai`` row carrying ``model`` / + ``temperature``, for instance) and those are stored and handed back to the + provider verbatim. + """ + + model_config = ConfigDict(extra="allow") + + +class OpenAICredentials(ProviderCredentialsBase): + api_key: str = Field(description="OpenAI API key") + + +class LangfuseCredentials(ProviderCredentialsBase): + secret_key: str = Field(description="Langfuse secret key") + public_key: str = Field(description="Langfuse public key") + host: str = Field(description="Langfuse host URL") + + +class GoogleAIStudioCredentials(ProviderCredentialsBase): + api_key: str = Field(description="Google AI Studio API key") + + +class SarvamAICredentials(ProviderCredentialsBase): + api_key: str = Field(description="SarvamAI API key") + + +class ElevenLabsCredentials(ProviderCredentialsBase): + api_key: str = Field(description="ElevenLabs API key") + + +class AnthropicCredentials(ProviderCredentialsBase): + api_key: str = Field(description="Anthropic API key") + + +class GoogleCredentials(ProviderCredentialsBase): + """Google Vertex AI (BYOK). + + Only ``api_key`` is required; each Vertex-specific field falls back to its + platform default in settings when omitted. + """ + + api_key: str = Field(description="Vertex AI API key") + project_id: str | None = Field(default=None, description="GCP project ID") + location: str | None = Field( + default=None, description="Vertex AI region, e.g. 'us-central1'" + ) + sa_key: dict[str, JsonValue] | None = Field( + default=None, description="GCP service-account key JSON" + ) + gcs_bucket: str | None = Field( + default=None, description="GCS bucket used to stage audio for Vertex" + ) + + +class WebhookSecretCredentials(ProviderCredentialsBase): + webhook_secret: str = Field( + description="Shared secret used to HMAC-sign outgoing webhooks" + ) + + +class ProxyCredentials(ProviderCredentialsBase): + api_key: str = Field(description="Proxy API key") + + +ProviderCredentials = Annotated[ + OpenAICredentials + | LangfuseCredentials + | GoogleAIStudioCredentials + | SarvamAICredentials + | ElevenLabsCredentials + | AnthropicCredentials + | GoogleCredentials + | WebhookSecretCredentials + | ProxyCredentials, + Field( + description=( + "Credential payload. The accepted shape depends on the provider it " + "is keyed by — see the per-provider schemas." + ) + ), +] + + @dataclass class ProviderConfig: - """Configuration for a provider including its required credential fields.""" + """Configuration for a provider: its credential schema and which of those + fields are secret.""" - required_fields: List[str] - sensitive_fields: List[str] = field(default_factory=list) + model: type[ProviderCredentialsBase] + sensitive_fields: list[str] = field(default_factory=list) + + @property + def required_fields(self) -> list[str]: + """Credential keys the provider cannot be used without, derived from + the schema so the two can never drift.""" + return [ + name + for name, model_field in self.model.model_fields.items() + if model_field.is_required() + ] # Provider configurations -PROVIDER_CONFIGS: Dict[Provider, ProviderConfig] = { +PROVIDER_CONFIGS: dict[Provider, ProviderConfig] = { Provider.OPENAI: ProviderConfig( - required_fields=["api_key"], sensitive_fields=["api_key"] + model=OpenAICredentials, sensitive_fields=["api_key"] ), Provider.LANGFUSE: ProviderConfig( - required_fields=["secret_key", "public_key", "host"], + model=LangfuseCredentials, sensitive_fields=["secret_key"], ), Provider.GOOGLE_AISTUDIO: ProviderConfig( - required_fields=["api_key"], sensitive_fields=["api_key"] + model=GoogleAIStudioCredentials, sensitive_fields=["api_key"] ), Provider.SARVAMAI: ProviderConfig( - required_fields=["api_key"], sensitive_fields=["api_key"] + model=SarvamAICredentials, sensitive_fields=["api_key"] ), Provider.ELEVENLABS: ProviderConfig( - required_fields=["api_key"], sensitive_fields=["api_key"] + model=ElevenLabsCredentials, sensitive_fields=["api_key"] ), Provider.ANTHROPIC: ProviderConfig( - required_fields=["api_key"], sensitive_fields=["api_key"] + model=AnthropicCredentials, sensitive_fields=["api_key"] ), Provider.GOOGLE: ProviderConfig( - required_fields=[ - "api_key", - ], + model=GoogleCredentials, sensitive_fields=["api_key"], ), Provider.WEBHOOK_SECRET: ProviderConfig( - required_fields=["webhook_secret"], sensitive_fields=["webhook_secret"] + model=WebhookSecretCredentials, sensitive_fields=["webhook_secret"] ), Provider.PROXY: ProviderConfig( - required_fields=["api_key"], sensitive_fields=["api_key"] + model=ProxyCredentials, sensitive_fields=["api_key"] ), } @@ -78,7 +179,7 @@ def validate_provider(provider: str) -> Provider: """ try: return Provider(provider.lower()) - except ValueError: + except (AttributeError, ValueError): # non-string keys reach here too supported = ", ".join(p.value for p in Provider) logger.warning( f"[validate_provider] Unsupported provider | provider: {provider}, supported_providers: {supported}" @@ -88,38 +189,70 @@ def validate_provider(provider: str) -> Provider: ) -def validate_provider_credentials(provider: str, credentials: Dict[str, str]) -> None: - """Validate that the credentials contain all required fields for the provider. +def validate_provider_credentials(provider: str, credentials: dict[str, Any]) -> None: + """Validate a credential payload against the provider's schema. Args: provider: The provider name to validate credentials for credentials: Dictionary containing the provider credentials Raises: - ValueError: If required fields are missing from the credentials + ValueError: If required fields are missing or a field has the wrong type + """ + parse_provider_credentials(provider, credentials) + + +def parse_provider_credentials( + provider: str, credentials: Any +) -> ProviderCredentialsBase: + """Validate a raw credential payload and return it as the provider's model. + + Unknown keys survive the round trip (see :class:`ProviderCredentialsBase`). + + Raises: + ValueError: If the provider is unsupported, the payload is not an + object, a required field is missing, or a field has the wrong type """ provider_enum = validate_provider(provider) - required_fields = PROVIDER_CONFIGS[provider_enum].required_fields + if not isinstance(credentials, dict): + raise ValueError(f"Value for provider '{provider}' must be an object/dict.") + + config = PROVIDER_CONFIGS[provider_enum] + + # Checked up front so the message stays field-oriented; Pydantic would + # otherwise report one "Field required" error per missing key. if missing_fields := [ - field for field in required_fields if field not in credentials + name for name in config.required_fields if name not in credentials ]: logger.warning( - f"[validate_provider_credentials] Missing required fields | provider: {provider}, missing_fields: {', '.join(missing_fields)}" + f"[parse_provider_credentials] Missing required fields | provider: {provider}, missing_fields: {', '.join(missing_fields)}" ) raise ValueError( f"Missing required fields for {provider}: {', '.join(missing_fields)}" ) + try: + return config.model.model_validate(credentials) + except ValidationError as e: + invalid_fields = ", ".join( + f"{'.'.join(str(loc) for loc in err['loc'])}: {err['msg']}" + for err in e.errors() + ) + logger.warning( + f"[parse_provider_credentials] Invalid credential fields | provider: {provider}, errors: {invalid_fields}" + ) + raise ValueError(f"Invalid credentials for {provider}: {invalid_fields}") + -def get_supported_providers() -> List[str]: +def get_supported_providers() -> list[str]: """Return a list of all supported provider names.""" return [p.value for p in Provider] def mask_credential_fields( - provider: str, credentials: Dict[str, Any] -) -> Dict[str, Any]: + provider: str, credentials: CredentialPayload +) -> CredentialPayload: """Mask sensitive fields in a credential dict for the given provider. Non-sensitive fields (e.g., langfuse `public_key`, `host`) are returned as-is. diff --git a/backend/app/crud/credentials.py b/backend/app/crud/credentials.py index cc070852b..eb8aa2827 100644 --- a/backend/app/crud/credentials.py +++ b/backend/app/crud/credentials.py @@ -26,7 +26,7 @@ def set_creds_for_org( ) raise HTTPException(400, "No credentials provided") - for provider, credentials in creds_add.credential.items(): + for provider, credentials in creds_add.credential_payloads().items(): try: validate_provider_credentials(provider, credentials) except ValueError as e: @@ -189,21 +189,14 @@ def update_creds_for_org( if not creds_in.provider or not creds_in.credential: raise ValueError("Provider and credential must be provided") - # Auto-unwrap nested format: {"google": {"api_key": "..."}} -> {"api_key": "..."} - # so the same payload shape works for both create and update. - credential_data = creds_in.credential - if ( - isinstance(credential_data, dict) - and creds_in.provider in credential_data - and isinstance(credential_data[creds_in.provider], dict) - ): - credential_data = credential_data[creds_in.provider] + provider = creds_in.provider.value + credential_data = creds_in.credential_payload() try: - validate_provider_credentials(creds_in.provider, credential_data) + validate_provider_credentials(provider, credential_data) except ValueError as e: logger.warning( - f"[update_creds_for_org] Validation error | organization_id: {org_id}, project_id: {project_id}, provider: {creds_in.provider}, error: {str(e)}" + f"[update_creds_for_org] Validation error | organization_id: {org_id}, project_id: {project_id}, provider: {provider}, error: {str(e)}" ) raise HTTPException(status_code=400, detail=str(e)) @@ -212,7 +205,7 @@ def update_creds_for_org( statement = select(Credential).where( Credential.organization_id == org_id, - Credential.provider == creds_in.provider, + Credential.provider == provider, Credential.is_active.is_(True), Credential.project_id == project_id, ) @@ -223,7 +216,7 @@ def update_creds_for_org( organization_id=org_id, project_id=project_id, is_active=creds_in.is_active if creds_in.is_active is not None else True, - provider=creds_in.provider, + provider=provider, credential=encrypted_credentials, inserted_at=now(), updated_at=now(), @@ -231,9 +224,6 @@ def update_creds_for_org( session.add(creds) session.commit() session.refresh(creds) - logger.info( - f"[update_creds_for_org] Created new credentials | organization_id {org_id}, provider {creds_in.provider}, project_id {project_id}" - ) return [creds] creds.credential = encrypted_credentials @@ -241,9 +231,6 @@ def update_creds_for_org( session.add(creds) session.commit() session.refresh(creds) - logger.info( - f"[update_creds_for_org] Successfully updated credentials | organization_id {org_id}, provider {creds_in.provider}, project_id {project_id}" - ) return [creds] @@ -328,6 +315,3 @@ def remove_creds_for_org(*, session: Session, org_id: int, project_id: int) -> N detail="Failed to delete all credentials", ) session.commit() - logger.info( - f"[remove_creds_for_org] Successfully deleted {rows_deleted} credential(s) | organization_id {org_id}, project_id {project_id}" - ) diff --git a/backend/app/crud/onboarding.py b/backend/app/crud/onboarding.py index e754414d0..2370c9674 100644 --- a/backend/app/crud/onboarding.py +++ b/backend/app/crud/onboarding.py @@ -120,16 +120,17 @@ def onboard_project( if onboard_in.credentials: for item in onboard_in.credentials: - (provider_str,) = item.keys() - values = item[provider_str] + ((provider, payload),) = item.items() - encrypted_credentials = encrypt_credentials(values) + encrypted_credentials = encrypt_credentials( + payload.model_dump(exclude_unset=True) + ) cred_row = Credential( organization_id=organization.id, project_id=project.id, is_active=True, - provider=provider_str, + provider=provider.value, credential=encrypted_credentials, ) session.add(cred_row) @@ -137,13 +138,6 @@ def onboard_project( created_credentials.append(cred_row) session.commit() - cred_ids = [c.id for c in created_credentials] - - logger.info( - "[onboard_project] Onboarding completed successfully. " - f"org_id={organization.id}, project_id={project.id}, user_id={user.id}, " - f"cred_ids={cred_ids}" - ) return OnboardingResponse( organization_id=organization.id, diff --git a/backend/app/models/credentials.py b/backend/app/models/credentials.py index e6741b2e9..de363f77f 100644 --- a/backend/app/models/credentials.py +++ b/backend/app/models/credentials.py @@ -2,9 +2,16 @@ from typing import Any import sqlalchemy as sa +from pydantic import field_validator, model_validator from sqlmodel import Field, Relationship, SQLModel -from app.core.providers import mask_credential_fields +from app.core.providers import ( + CredentialPayload, + Provider, + ProviderCredentials, + mask_credential_fields, + parse_provider_credentials, +) from app.core.util import now from app.models.organization import Organization from app.models.project import Project @@ -38,32 +45,76 @@ class CredsBase(SQLModel): class CredsCreate(SQLModel): """Create new credentials for an organization. - The credential field should be a dictionary mapping provider names to their credentials. + The credential field maps each provider to its own credential payload. Example: {"openai": {"api_key": "..."}, "langfuse": {"public_key": "..."}} """ is_active: bool = True - credential: dict[str, Any] = Field( + credential: dict[Provider, ProviderCredentials] | None = Field( default=None, - description="Dictionary mapping provider names to their credentials", + description="Credential payload per provider, keyed by provider name", ) + @field_validator("credential", mode="before") + @classmethod + def _parse_credential(cls, value: Any) -> Any: + if not isinstance(value, dict): + return value + + return { + provider: parse_provider_credentials(provider, payload) + for provider, payload in value.items() + } + + def credential_payloads(self) -> dict[str, dict[str, Any]]: + """Provider name -> credential dict, exactly as submitted.""" + return { + provider.value: payload.model_dump(exclude_unset=True) + for provider, payload in (self.credential or {}).items() + } + class CredsUpdate(SQLModel): """Update credentials for an organization. Can update a specific provider's credentials or add a new provider. """ - provider: str = Field( + provider: Provider = Field( description="Name of the provider to update/add credentials for" ) - credential: dict[str, Any] = Field( + credential: ProviderCredentials = Field( description="Credentials for the specified provider", ) is_active: bool | None = Field( default=None, description="Whether the credentials are active" ) + @model_validator(mode="before") + @classmethod + def _parse_credential(cls, data: Any) -> Any: + if not isinstance(data, dict): + return data + + provider = data.get("provider") + credential = data.get("credential") + if not isinstance(provider, str) or credential is None: + return data + + provider_key = provider.value if isinstance(provider, Provider) else provider + if isinstance(credential, dict) and isinstance( + credential.get(provider_key), dict + ): + credential = credential[provider_key] + + return { + **data, + "credential": parse_provider_credentials(provider_key, credential), + } + + def credential_payload(self) -> dict[str, Any]: + """Credential dict for `provider`, exactly as submitted.""" + return self.credential.model_dump(exclude_unset=True) + class Credential(CredsBase, table=True): """Database model for storing provider credentials. @@ -132,7 +183,7 @@ def to_public(self, mask: bool = True) -> "CredsPublic": organization_id=self.organization_id, project_id=self.project_id, is_active=self.is_active, - provider=self.provider, + provider=Provider(self.provider), credential=decrypted, inserted_at=self.inserted_at, updated_at=self.updated_at, @@ -143,7 +194,10 @@ class CredsPublic(CredsBase): """Public representation of credentials, excluding sensitive information.""" id: int - provider: str - credential: dict[str, Any] | None = None + provider: Provider + credential: CredentialPayload | None = Field( + default=None, + description="Provider credential with its sensitive fields masked", + ) inserted_at: datetime updated_at: datetime diff --git a/backend/app/models/onboarding.py b/backend/app/models/onboarding.py index 942d36134..8ab1dab7a 100644 --- a/backend/app/models/onboarding.py +++ b/backend/app/models/onboarding.py @@ -5,7 +5,11 @@ from sqlmodel import SQLModel, Field from pydantic import EmailStr, model_validator, field_validator -from app.core.providers import validate_provider, validate_provider_credentials +from app.core.providers import ( + Provider, + ProviderCredentials, + parse_provider_credentials, +) class OnboardingRequest(SQLModel): @@ -52,9 +56,13 @@ class OnboardingRequest(SQLModel): min_length=3, max_length=50, ) - credentials: list[dict[str, Any]] | None = Field( + credentials: list[dict[Provider, ProviderCredentials]] | None = Field( default=None, - description="Optional credential(s) to link with the project", + description=( + "Optional credential(s) to link with the project. Each entry maps " + "exactly one provider to its credential payload, e.g. " + "{'openai': {'api_key': '...'}}" + ), ) @staticmethod @@ -82,40 +90,29 @@ def set_defaults(self): self.password = secrets.token_urlsafe(12) return self - @field_validator("credentials") + @field_validator("credentials", mode="before") @classmethod - def _validate_credential_list(cls, v: list[dict[str, dict[str, str]]] | None): - if v is None: - return v - - if not isinstance(v, list): - raise TypeError( - "Credential must be a list of single-key dicts (e.g., {'openai': {...}})." - ) - - for item in v: - if not isinstance(item, dict): - raise TypeError( - "Credential must be a dict with a single provider key like {'openai': {...}}." - ) + def _parse_credential_list(cls, value: Any) -> Any: + # Ill-typed containers are handed back untouched so Pydantic reports + # the shape mismatch ("Input should be a valid list/dictionary"). + if not isinstance(value, list) or any( + not isinstance(item, dict) for item in value + ): + return value + + parsed: list[dict[str, Any]] = [] + for item in value: if len(item) != 1: raise ValueError( "Credential must have exactly one provider key like {'openai': {...}}." ) - (provider_key,) = item.keys() - values = item[provider_key] - - validate_provider(provider_key) - - if not isinstance(values, dict): - raise ValueError( - f"Value for provider '{provider_key}' must be an object/dict." - ) - - validate_provider_credentials(provider_key, values) + (provider,) = item.keys() + parsed.append( + {provider: parse_provider_credentials(provider, item[provider])} + ) - return v + return parsed class OnboardingResponse(SQLModel): diff --git a/backend/app/models/project.py b/backend/app/models/project.py index f5902402a..67b5e88fa 100644 --- a/backend/app/models/project.py +++ b/backend/app/models/project.py @@ -2,6 +2,7 @@ from typing import TYPE_CHECKING, Any, Optional from uuid import UUID, uuid4 +from pydantic import JsonValue from sqlalchemy import Column from sqlalchemy.dialects.postgresql import JSONB from sqlmodel import Field, Relationship, SQLModel, UniqueConstraint @@ -134,7 +135,7 @@ class Project(ProjectBase, table=True): class ProjectPublic(ProjectBase): id: int organization_id: int - settings: dict[str, Any] = {} + settings: dict[str, JsonValue] = {} inserted_at: datetime updated_at: datetime diff --git a/backend/app/tests/api/routes/test_creds.py b/backend/app/tests/api/routes/test_creds.py index c5fd9c1ea..3aa6c02d0 100644 --- a/backend/app/tests/api/routes/test_creds.py +++ b/backend/app/tests/api/routes/test_creds.py @@ -482,8 +482,11 @@ def test_update_credential_empty_credential( headers={"X-API-KEY": user_api_key.key}, ) - # Should still update but with empty credential - assert response.status_code in [200, 400] # Depends on implementation + assert response.status_code == 422 + assert any( + "Missing required fields for openai" in error["message"] + for error in response.json()["errors"] + ) def test_read_credentials_when_none_exist( diff --git a/backend/app/tests/crud/test_credentials.py b/backend/app/tests/crud/test_credentials.py index 8b3269ca6..df612727c 100644 --- a/backend/app/tests/crud/test_credentials.py +++ b/backend/app/tests/crud/test_credentials.py @@ -1,4 +1,5 @@ import pytest +from pydantic import ValidationError from sqlmodel import Session from app.crud import ( @@ -135,7 +136,9 @@ def test_set_credentials_for_google_with_sa_key(db: Session) -> None: def test_get_provider_credential(db: Session) -> None: """Test retrieving credentials for a specific provider.""" credentials_create = test_credential_data(db) - original_api_key = credentials_create.credential[Provider.OPENAI.value]["api_key"] + original_api_key = credentials_create.credential_payloads()[Provider.OPENAI.value][ + "api_key" + ] project = create_test_project(db) set_creds_for_org( @@ -250,28 +253,17 @@ def test_remove_creds_for_org(db: Session) -> None: assert creds == [] -def test_invalid_provider(db: Session) -> None: - """Test handling of invalid provider names.""" - from app.core.exception_handlers import HTTPException - - project = create_test_project(db) - +def test_invalid_provider() -> None: + """Unsupported providers are rejected while building the request model.""" credentials_data = {"invalid_provider": {"api_key": "test-key"}} - credentials_create = CredsCreate( - is_active=True, - credential=credentials_data, - ) - with pytest.raises(HTTPException) as exc_info: - set_creds_for_org( - session=db, - creds_add=credentials_create, - organization_id=project.organization_id, - project_id=project.id, + with pytest.raises(ValidationError) as exc_info: + CredsCreate( + is_active=True, + credential=credentials_data, ) - assert exc_info.value.status_code == 400 - assert "Unsupported provider" in exc_info.value.detail + assert "Unsupported provider" in str(exc_info.value) def test_duplicate_provider_credentials(db: Session) -> None: @@ -304,8 +296,6 @@ def test_duplicate_provider_credentials(db: Session) -> None: def test_langfuse_credential_validation(db: Session) -> None: """Test validation of Langfuse credentials structure.""" - from app.core.exception_handlers import HTTPException - project = create_test_project(db) # Test with missing required fields @@ -316,21 +306,14 @@ def test_langfuse_credential_validation(db: Session) -> None: # Missing host } } - credentials_create = CredsCreate( - is_active=True, - credential=invalid_credentials, - ) - with pytest.raises(HTTPException) as exc_info: - set_creds_for_org( - session=db, - creds_add=credentials_create, - organization_id=project.organization_id, - project_id=project.id, + with pytest.raises(ValidationError) as exc_info: + CredsCreate( + is_active=True, + credential=invalid_credentials, ) - assert exc_info.value.status_code == 400 - assert "Missing required fields for langfuse" in exc_info.value.detail + assert "Missing required fields for langfuse" in str(exc_info.value) valid_credentials = { "langfuse": { diff --git a/docs/wiki/modules/platform.md b/docs/wiki/modules/platform.md index 0f34e5d87..e07dfa086 100644 --- a/docs/wiki/modules/platform.md +++ b/docs/wiki/modules/platform.md @@ -14,3 +14,26 @@ All paths relative to `backend/app/`. | Model config | `api/routes/model_config.py` | `model_config` (`models/model_config.py`) | `crud/model_config.py` | | Cron | `api/routes/cron.py` | — | triggers batch polling (`crud/evaluations/cron.py`) | | Jobs | — | `job` (`models/job.py`), `batch_job` (`models/batch_job.py`) | `crud/jobs.py`, `crud/job/`, `services/job_monitoring.py` | + +## Credential contracts + +`core/providers.py` holds one Pydantic model per `Provider` (`OpenAICredentials`, +`LangfuseCredentials`, `GoogleCredentials`, …) and `PROVIDER_CONFIGS` maps each +provider to its model plus its sensitive (masked) fields. Those models are the +single source of truth: + +- `ProviderConfig.required_fields` is derived from the model, so a schema change + and the validation gate can't drift apart. +- Request models (`CredsCreate.credential`, `CredsUpdate.credential`, + `OnboardingRequest.credentials`) are keyed by `Provider` and typed as the union + of those models, so the shape lands in the OpenAPI schema instead of + `dict[str, Any]` — adding or changing a required field is a visible contract + change. +- Payloads allow extra keys (provider passthrough, e.g. `model` on an `openai` + row); they round-trip via `model_dump(exclude_unset=True)`. +- Responses (`CredsPublic.credential`, `GET /credentials/provider/{provider}`) + are masked, so they are typed `CredentialPayload` (`dict[str, JsonValue]`) + rather than a provider model. + +Adding a provider = add the enum member, the model, and its `PROVIDER_CONFIGS` +entry; nothing else needs to know. diff --git a/docs/wiki/modules/tenancy.md b/docs/wiki/modules/tenancy.md index 178184590..9da23b418 100644 --- a/docs/wiki/modules/tenancy.md +++ b/docs/wiki/modules/tenancy.md @@ -24,3 +24,12 @@ All paths relative to `backend/app/`. ## Services / CRUD - `crud/user.py`, `crud/organization.py`, `crud/project.py`, `crud/user_project.py`, `crud/api_key.py`, `crud/auth.py`, `crud/onboarding.py` - `services/auth.py` + +## API contracts +- `OnboardingRequest.credentials` is a list of single-entry `{Provider: payload}` + maps typed against the per-provider credential models in `core/providers.py` + (see `modules/platform.md`). Unsupported providers, non-object payloads and + missing required fields are rejected at the request boundary (422). +- `ProjectPublic.settings` is `dict[str, JsonValue]` — free-form JSONB whose + writable keys are defined by `ProjectSettingsUpdate` (`PATCH + /projects/settings`), currently just `tracing`.