"""凭据存储与提供商探测的接口 + SQLAlchemy/网关实现(ARCH §4.7)。 路由依赖这里的 **接口**(Protocol),测试注入内存替身;运行时用 SQLAlchemy/网关实现。 单用户原型:`owner_id` 用固定 stub(见 `STUB_OWNER_ID`),多租户化时改为按认证主体取。 """ from __future__ import annotations import uuid from dataclasses import dataclass from typing import Protocol from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from ww_db.models import ProviderCredential, TierRouting from ww_llm_gateway.adapters.base import Capabilities # 单用户 stub owner(与网关 Scope.user_id 的 stub 约定一致:UUID(int=1))。 # 多租户化时此常量由认证主体替换(见 ARCH §4.7 隔离)。 STUB_OWNER_ID = uuid.UUID(int=1) # 凭据认证类型(`provider_credentials.auth_type`,见 C2 扩 K1.1)。 AUTH_TYPE_API_KEY = "api_key" AUTH_TYPE_OAUTH = "oauth" @dataclass(frozen=True) class StoredCredential: """存储层视图:含密文,绝不出 API 边界(路由仅取 provider 并掩码)。 一行二选一:`auth_type="api_key"` → `api_key_enc` 有值、`oauth_enc=None`; `auth_type="oauth"`(Kimi Code device-flow,K1.3)→ `oauth_enc` 有值、`api_key_enc=None` (持 Fernet 加密的 `{access_token,refresh_token,expires_at}` JSON 包)。 """ provider: str api_key_enc: bytes | None auth_type: str = AUTH_TYPE_API_KEY oauth_enc: bytes | None = None @dataclass(frozen=True) class StoredRouting: tier: str provider: str model: str fallback: list[str] class CredentialStore(Protocol): """凭据 + 档位路由的读写接口(按 owner_id 隔离)。""" async def list_credentials(self, owner_id: uuid.UUID) -> list[StoredCredential]: ... async def list_routing(self) -> list[StoredRouting]: ... async def get_credential( self, owner_id: uuid.UUID, provider: str ) -> StoredCredential | None: ... async def upsert_credential( self, owner_id: uuid.UUID, provider: str, api_key_enc: bytes ) -> None: ... async def upsert_oauth_credential( self, owner_id: uuid.UUID, provider: str, oauth_enc: bytes ) -> None: ... async def delete_credential(self, owner_id: uuid.UUID, provider: str) -> bool: ... async def upsert_routing(self, routing: StoredRouting) -> None: ... class ProviderProbe(Protocol): """最小连通探测:验证 Key + 返回能力矩阵。测试注入假探测,绝不联网。""" async def probe(self, owner_id: uuid.UUID, provider: str) -> Capabilities: ... class SqlCredentialStore: """SQLAlchemy 实现:写 `provider_credentials` / `tier_routing`,幂等 upsert。""" def __init__(self, session: AsyncSession) -> None: self._session = session async def list_credentials(self, owner_id: uuid.UUID) -> list[StoredCredential]: rows = ( await self._session.execute( select(ProviderCredential).where(ProviderCredential.owner_id == owner_id) ) ).scalars() return [ StoredCredential( provider=r.provider, api_key_enc=r.api_key_enc, auth_type=r.auth_type, oauth_enc=r.oauth_enc, ) for r in rows ] async def list_routing(self) -> list[StoredRouting]: rows = (await self._session.execute(select(TierRouting))).scalars() return [ StoredRouting( tier=r.tier, provider=r.provider, model=r.model, fallback=list(r.fallback) ) for r in rows ] async def get_credential(self, owner_id: uuid.UUID, provider: str) -> StoredCredential | None: row = ( await self._session.execute( select(ProviderCredential).where( ProviderCredential.owner_id == owner_id, ProviderCredential.provider == provider, ) ) ).scalar_one_or_none() if row is None: return None return StoredCredential( provider=row.provider, api_key_enc=row.api_key_enc, auth_type=row.auth_type, oauth_enc=row.oauth_enc, ) async def upsert_credential( self, owner_id: uuid.UUID, provider: str, api_key_enc: bytes ) -> None: # 显式 read-modify-write:唯一约束含可空 project_id,PG ON CONFLICT # 在 NULL 上不去重(NULLS DISTINCT),故不用 on_conflict。 existing = ( await self._session.execute( select(ProviderCredential).where( ProviderCredential.owner_id == owner_id, ProviderCredential.project_id.is_(None), ProviderCredential.provider == provider, ) ) ).scalar_one_or_none() if existing is None: self._session.add( ProviderCredential( owner_id=owner_id, project_id=None, provider=provider, api_key_enc=api_key_enc, auth_type=AUTH_TYPE_API_KEY, oauth_enc=None, ) ) else: existing.api_key_enc = api_key_enc existing.auth_type = AUTH_TYPE_API_KEY existing.oauth_enc = None await self._session.commit() async def upsert_oauth_credential( self, owner_id: uuid.UUID, provider: str, oauth_enc: bytes ) -> None: """写/更新 OAuth 凭据行(Kimi Code device-flow,K1.3)。 `auth_type="oauth"`、`oauth_enc=`、`api_key_enc=None`。 显式 read-modify-write(同 `upsert_credential`:含可空 project_id 的唯一约束不能用 PG `ON CONFLICT`,见 memory/gotchas)。明文 token 绝不进此层(已加密)。 """ existing = ( await self._session.execute( select(ProviderCredential).where( ProviderCredential.owner_id == owner_id, ProviderCredential.project_id.is_(None), ProviderCredential.provider == provider, ) ) ).scalar_one_or_none() if existing is None: self._session.add( ProviderCredential( owner_id=owner_id, project_id=None, provider=provider, api_key_enc=None, auth_type=AUTH_TYPE_OAUTH, oauth_enc=oauth_enc, ) ) else: existing.api_key_enc = None existing.auth_type = AUTH_TYPE_OAUTH existing.oauth_enc = oauth_enc await self._session.commit() async def delete_credential(self, owner_id: uuid.UUID, provider: str) -> bool: """删除凭据行(OAuth disconnect / 撤销)。返回是否删到行。""" existing = ( await self._session.execute( select(ProviderCredential).where( ProviderCredential.owner_id == owner_id, ProviderCredential.project_id.is_(None), ProviderCredential.provider == provider, ) ) ).scalar_one_or_none() if existing is None: return False await self._session.delete(existing) await self._session.commit() return True async def upsert_routing(self, routing: StoredRouting) -> None: existing = ( await self._session.execute( select(TierRouting).where( TierRouting.project_id.is_(None), TierRouting.tier == routing.tier, ) ) ).scalar_one_or_none() if existing is None: self._session.add( TierRouting( project_id=None, tier=routing.tier, provider=routing.provider, model=routing.model, fallback=routing.fallback, ) ) else: existing.provider = routing.provider existing.model = routing.model existing.fallback = routing.fallback await self._session.commit()