from __future__ import annotations import sqlite3 import unicodedata from datetime import datetime from urllib.parse import urlsplit from backend.database.connection import Database from backend.errors import BusinessError from backend.features.accounts.models import ( ModelPoolItem, ModelPoolRecord, ModelRuntimeConfig, Principal, ) from backend.features.accounts.service import now_utc from backend.security import SecretCipher def _record(row: sqlite3.Row) -> ModelPoolRecord: return ModelPoolRecord( id=int(row["id"]), display_name=str(row["display_name"]), display_name_key=str(row["display_name_key"]), base_url=str(row["base_url"]), model_identifier=str(row["model_identifier"]), encrypted_api_key=str(row["encrypted_api_key"]), created_at=datetime.fromisoformat(str(row["created_at"])), updated_at=datetime.fromisoformat(str(row["updated_at"])), updated_by=int(row["updated_by"]), ) class ModelPoolRepository: @staticmethod def count(connection: sqlite3.Connection) -> int: return int(connection.execute("SELECT COUNT(*) FROM llm_models").fetchone()[0]) @staticmethod def create( connection: sqlite3.Connection, display_name: str, display_name_key: str, base_url: str, model_identifier: str, encrypted_api_key: str, actor_id: int, now: datetime, ) -> ModelPoolRecord: timestamp = now.isoformat(timespec="seconds") cursor = connection.execute( """ INSERT INTO llm_models ( display_name, display_name_key, base_url, model_identifier, encrypted_api_key, created_at, updated_at, updated_by ) VALUES (?, ?, ?, ?, ?, ?, ?, ?) """, ( display_name, display_name_key, base_url, model_identifier, encrypted_api_key, timestamp, timestamp, actor_id, ), ) row = connection.execute( "SELECT * FROM llm_models WHERE id = ?", (cursor.lastrowid,) ).fetchone() return _record(row) @staticmethod def get(connection: sqlite3.Connection, model_id: int) -> ModelPoolRecord | None: row = connection.execute( "SELECT * FROM llm_models WHERE id = ?", (model_id,) ).fetchone() return _record(row) if row else None @staticmethod def list(connection: sqlite3.Connection) -> tuple[ModelPoolRecord, ...]: return tuple( _record(row) for row in connection.execute("SELECT * FROM llm_models ORDER BY id") ) @staticmethod def update( connection: sqlite3.Connection, model_id: int, display_name: str, display_name_key: str, base_url: str, model_identifier: str, encrypted_api_key: str, actor_id: int, now: datetime, ) -> ModelPoolRecord | None: connection.execute( """ UPDATE llm_models SET display_name = ?, display_name_key = ?, base_url = ?, model_identifier = ?, encrypted_api_key = ?, updated_at = ?, updated_by = ? WHERE id = ? """, ( display_name, display_name_key, base_url, model_identifier, encrypted_api_key, now.isoformat(timespec="seconds"), actor_id, model_id, ), ) return ModelPoolRepository.get(connection, model_id) @staticmethod def delete(connection: sqlite3.Connection, model_id: int) -> None: connection.execute("DELETE FROM llm_models WHERE id = ?", (model_id,)) @staticmethod def selection(connection: sqlite3.Connection) -> tuple[int | None, int | None]: row = connection.execute( "SELECT primary_model_id, fallback_model_id FROM llm_configuration WHERE id = 1" ).fetchone() return row["primary_model_id"], row["fallback_model_id"] @staticmethod def save_selection( connection: sqlite3.Connection, primary_id: int, fallback_id: int | None, actor_id: int, now: datetime, ) -> None: connection.execute( """ UPDATE llm_configuration SET primary_model_id = ?, fallback_model_id = ?, updated_at = ?, updated_by = ? WHERE id = 1 """, ( primary_id, fallback_id, now.isoformat(timespec="seconds"), actor_id, ), ) class ModelPoolService: MAX_MODELS = 20 def __init__( self, database: Database, repository: ModelPoolRepository, cipher: SecretCipher, ) -> None: self._database = database self._repository = repository self._cipher = cipher def list(self) -> tuple[ModelPoolItem, ...]: with self._database.read() as connection: records = self._repository.list(connection) primary_id, fallback_id = self._repository.selection(connection) return tuple( self._to_item(record, primary_id, fallback_id) for record in records ) def get_item(self, model_id: int) -> ModelPoolItem: with self._database.read() as connection: record = self._repository.get(connection, model_id) primary_id, fallback_id = self._repository.selection(connection) if record is None: raise BusinessError("model_not_found", "模型不存在。") return self._to_item(record, primary_id, fallback_id) def create( self, actor: Principal, display_name: str, base_url: str, model_identifier: str, api_key: str, ) -> ModelPoolItem: name, key, url, identifier = self._validate( display_name, base_url, model_identifier ) cleaned_key = self._validate_api_key(api_key) now = now_utc() try: with self._database.transaction() as connection: if self._repository.count(connection) >= self.MAX_MODELS: raise BusinessError("model_pool_full", "模型池最多可添加20个模型。") record = self._repository.create( connection, name, key, url, identifier, self._cipher.encrypt(cleaned_key), actor.user.id, now, ) primary_id, _ = self._repository.selection(connection) if primary_id is None: self._repository.save_selection( connection, record.id, None, actor.user.id, now ) return self.get_item(record.id) except sqlite3.IntegrityError as exc: raise BusinessError("model_name_taken", "模型显示名称已存在。") from exc def update( self, actor: Principal, model_id: int, display_name: str, base_url: str, model_identifier: str, api_key: str | None, ) -> ModelPoolItem: name, key, url, identifier = self._validate( display_name, base_url, model_identifier ) now = now_utc() try: with self._database.transaction() as connection: current = self._repository.get(connection, model_id) if current is None: raise BusinessError("model_not_found", "模型不存在。") encrypted_key = ( self._cipher.encrypt(self._validate_api_key(api_key)) if api_key is not None else current.encrypted_api_key ) updated = self._repository.update( connection, model_id, name, key, url, identifier, encrypted_key, actor.user.id, now, ) if updated is None: raise BusinessError("model_not_found", "模型不存在。") return self.get_item(updated.id) except sqlite3.IntegrityError as exc: raise BusinessError("model_name_taken", "模型显示名称已存在。") from exc def delete(self, model_id: int) -> None: with self._database.transaction() as connection: current = self._repository.get(connection, model_id) if current is None: raise BusinessError("model_not_found", "模型不存在。") primary_id, fallback_id = self._repository.selection(connection) if model_id in {primary_id, fallback_id}: raise BusinessError("model_in_use", "请先调整主模型或辅助模型后再删除。") self._repository.delete(connection, model_id) def select( self, actor: Principal, primary_id: int, fallback_id: int | None ) -> None: if fallback_id == primary_id: raise BusinessError("duplicate_model_role", "主模型和辅助模型不能相同。") with self._database.transaction() as connection: if self._repository.get(connection, primary_id) is None: raise BusinessError("model_not_found", "主模型不存在。") if fallback_id is not None and self._repository.get(connection, fallback_id) is None: raise BusinessError("model_not_found", "辅助模型不存在。") self._repository.save_selection( connection, primary_id, fallback_id, actor.user.id, now_utc() ) def runtime_config(self) -> ModelRuntimeConfig: with self._database.read() as connection: primary_id, fallback_id = self._repository.selection(connection) primary = self._repository.get(connection, primary_id) if primary_id else None fallback = self._repository.get(connection, fallback_id) if fallback_id else None if primary is None: raise BusinessError("model_not_configured", "智能解读服务尚未配置。") return ModelRuntimeConfig(primary=primary, fallback=fallback) def decrypt_api_key(self, record: ModelPoolRecord) -> str: return self._cipher.decrypt(record.encrypted_api_key) @staticmethod def _to_item( record: ModelPoolRecord, primary_id: int | None, fallback_id: int | None, ) -> ModelPoolItem: return ModelPoolItem( id=record.id, display_name=record.display_name, base_url=record.base_url, model_identifier=record.model_identifier, has_api_key=bool(record.encrypted_api_key), is_primary=record.id == primary_id, is_fallback=record.id == fallback_id, updated_at=record.updated_at, ) @staticmethod def _validate( display_name: str, base_url: str, model_identifier: str ) -> tuple[str, str, str, str]: name = unicodedata.normalize("NFKC", display_name.strip()) if not 1 <= len(name) <= 40: raise BusinessError("invalid_model", "模型显示名称应为1至40个字符。") url = base_url.strip().rstrip("/") parsed = urlsplit(url) if parsed.scheme not in {"http", "https"} or not parsed.netloc: raise BusinessError("invalid_model", "模型服务地址无效。") if parsed.username or parsed.password: raise BusinessError("invalid_model", "模型服务地址不能包含账号或密码。") identifier = model_identifier.strip() if not 1 <= len(identifier) <= 128: raise BusinessError("invalid_model", "模型标识应为1至128个字符。") return name, name.casefold(), url, identifier @staticmethod def _validate_api_key(api_key: str) -> str: cleaned = api_key.strip() if not 1 <= len(cleaned) <= 4096: raise BusinessError("invalid_model", "模型密钥无效。") return cleaned