diff --git a/backend/app/api/endpoints/ai_models.py b/backend/app/api/endpoints/ai_models.py index 0f88fd4..b2d3c34 100755 --- a/backend/app/api/endpoints/ai_models.py +++ b/backend/app/api/endpoints/ai_models.py @@ -1,8 +1,11 @@ """AI模型配置接口""" -from fastapi import APIRouter, Depends, HTTPException +import secrets + +from fastapi import APIRouter, Depends, Header, HTTPException from sqlalchemy import select, update from app.core.database import get_db +from app.core.config import settings from app.schemas import ApiResponse, AIModelCreateRequest, AIModelUpdateRequest, AIModelTestRequest from app.models import AIModelConfig from app.utils.crypto import encrypt, decrypt @@ -22,10 +25,15 @@ async def list_models(db=Depends(get_db)): @router.post("") async def create_model(req: AIModelCreateRequest, db=Depends(get_db)): if req.is_default: - await db.execute(update(AIModelConfig).values(is_default=0)) + await db.execute( + update(AIModelConfig) + .where(AIModelConfig.usage_scope == req.usage_scope) + .values(is_default=0) + ) model = AIModelConfig( model_name=req.model_name, provider=req.provider, + usage_scope=req.usage_scope, api_base_url=req.api_base_url, api_key_enc=encrypt(req.api_key) if req.api_key else None, model_version=req.model_version, @@ -47,8 +55,16 @@ async def update_model(model_id: int, req: AIModelUpdateRequest, db=Depends(get_ model = result.scalar_one_or_none() if not model: raise HTTPException(status_code=404, detail="模型不存在") - if req.is_default: - await db.execute(update(AIModelConfig).where(AIModelConfig.id != model_id).values(is_default=0)) + target_scope = req.usage_scope or model.usage_scope + if req.is_default or (req.usage_scope and model.is_default): + await db.execute( + update(AIModelConfig) + .where( + AIModelConfig.id != model_id, + AIModelConfig.usage_scope == target_scope, + ) + .values(is_default=0) + ) for field, val in req.model_dump(exclude_none=True).items(): if field == "api_key": model.api_key_enc = encrypt(val) if val else None @@ -59,6 +75,37 @@ async def update_model(model_id: int, req: AIModelUpdateRequest, db=Depends(get_ return ApiResponse(data=_format_model(model), message="更新成功") +@router.get("/runtime/digital-avatar") +async def get_digital_avatar_runtime_model( + x_avatar_config_token: str | None = Header(default=None), + db=Depends(get_db), +): + expected = settings.AVATAR_MODEL_CONFIG_TOKEN + if not expected: + raise HTTPException(status_code=503, detail="数字分身模型配置服务未启用") + if not x_avatar_config_token or not secrets.compare_digest(x_avatar_config_token, expected): + raise HTTPException(status_code=401, detail="无权读取数字分身模型配置") + + result = await db.execute( + select(AIModelConfig).where( + AIModelConfig.usage_scope == "digital_avatar", + AIModelConfig.is_default == 1, + AIModelConfig.is_enabled == 1, + ) + ) + model = result.scalar_one_or_none() + if not model: + raise HTTPException(status_code=404, detail="尚未配置启用的数字分身专用模型") + return ApiResponse(data={ + "api_base_url": model.api_base_url or "https://api.openai.com/v1", + "api_key": decrypt(model.api_key_enc) if model.api_key_enc else "", + "model": model.model_version or model.model_name, + "temperature": model.temperature, + "max_tokens": model.max_tokens, + "timeout_seconds": model.timeout_seconds, + }) + + @router.delete("/{model_id}") async def delete_model(model_id: int, db=Depends(get_db)): result = await db.execute(select(AIModelConfig).where(AIModelConfig.id == model_id)) @@ -79,6 +126,7 @@ async def test_model(req: AIModelTestRequest, db=Depends(get_db)): def _format_model(m: AIModelConfig) -> dict: return { "id": m.id, "model_name": m.model_name, "provider": m.provider, + "usage_scope": m.usage_scope, "api_base_url": m.api_base_url, "has_api_key": bool(m.api_key_enc), "model_version": m.model_version, "temperature": m.temperature, "max_tokens": m.max_tokens, "timeout_seconds": m.timeout_seconds, diff --git a/backend/app/core/config.py b/backend/app/core/config.py index 34844b5..a07b722 100755 --- a/backend/app/core/config.py +++ b/backend/app/core/config.py @@ -19,6 +19,7 @@ class Settings(BaseSettings): # 安全 SECRET_KEY: str = os.getenv("SECRET_KEY", "dev-secret-key-change-in-prod") AES_KEY: str = os.getenv("AES_KEY", "your-aes-key-32-chars-change-now!") + AVATAR_MODEL_CONFIG_TOKEN: str = os.getenv("AVATAR_MODEL_CONFIG_TOKEN", "") # 新闻平台 NEWS_PLATFORM_BASE_URL: str = os.getenv( diff --git a/backend/app/core/database.py b/backend/app/core/database.py index 6023207..ae3ccd6 100755 --- a/backend/app/core/database.py +++ b/backend/app/core/database.py @@ -1,6 +1,7 @@ """数据库连接管理""" import asyncio from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker +from sqlalchemy import text from sqlalchemy.orm import DeclarativeBase from app.core.config import settings from app.core.logger import logger @@ -64,6 +65,18 @@ async def init_db(): VirtualUser, UserPersonality, InteractionRecord, PendingReplyTask, TokenStat, AIModelConfig, SystemConfig, LoginLog ) + async with engine.begin() as conn: + result = await conn.execute(text( + "SELECT COUNT(*) FROM information_schema.COLUMNS " + "WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'ai_model_configs' " + "AND COLUMN_NAME = 'usage_scope'" + )) + if result.scalar_one() == 0: + await conn.execute(text( + "ALTER TABLE ai_model_configs ADD COLUMN usage_scope " + "VARCHAR(16) NOT NULL DEFAULT 'general' AFTER provider" + )) + logger.info("AI模型配置表已增加 usage_scope 字段") logger.info("✅ 数据库模型注册成功") logger.info("✅ 数据库初始化完成") diff --git a/backend/app/models/__init__.py b/backend/app/models/__init__.py index 75090ad..aa33965 100755 --- a/backend/app/models/__init__.py +++ b/backend/app/models/__init__.py @@ -122,6 +122,7 @@ class AIModelConfig(Base): id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) model_name: Mapped[str] = mapped_column(String(64), nullable=False) provider: Mapped[str] = mapped_column(String(32), nullable=False) + usage_scope: Mapped[str] = mapped_column(String(16), nullable=False, default="general") api_base_url: Mapped[str | None] = mapped_column(String(256)) api_key_enc: Mapped[str | None] = mapped_column(String(512)) model_version: Mapped[str | None] = mapped_column(String(64)) diff --git a/backend/app/schemas/__init__.py b/backend/app/schemas/__init__.py index 68eeab7..bc76534 100755 --- a/backend/app/schemas/__init__.py +++ b/backend/app/schemas/__init__.py @@ -154,6 +154,7 @@ class InteractionResponse(BaseModel): class AIModelCreateRequest(BaseModel): model_name: str = Field(..., min_length=1, max_length=64) provider: str = Field(..., pattern="^(openai|zhipu|wenxin|qianwen|local)$") + usage_scope: str = Field(default="general", pattern="^(general|digital_avatar)$") api_base_url: Optional[str] = None api_key: Optional[str] = None model_version: Optional[str] = None @@ -165,6 +166,8 @@ class AIModelCreateRequest(BaseModel): class AIModelUpdateRequest(BaseModel): model_name: Optional[str] = None + provider: Optional[str] = Field(None, pattern="^(openai|zhipu|wenxin|qianwen|local)$") + usage_scope: Optional[str] = Field(None, pattern="^(general|digital_avatar)$") api_base_url: Optional[str] = None api_key: Optional[str] = None model_version: Optional[str] = None @@ -179,6 +182,7 @@ class AIModelResponse(BaseModel): id: int model_name: str provider: str + usage_scope: str api_base_url: Optional[str] has_api_key: bool model_version: Optional[str] diff --git a/backend/app/services/ai_service.py b/backend/app/services/ai_service.py index a475950..0322385 100755 --- a/backend/app/services/ai_service.py +++ b/backend/app/services/ai_service.py @@ -28,7 +28,9 @@ class AIService: async def _get_default_model(self, db: AsyncSession) -> Optional[AIModelConfig]: result = await db.execute( select(AIModelConfig).where( - AIModelConfig.is_default == 1, AIModelConfig.is_enabled == 1 + AIModelConfig.usage_scope == "general", + AIModelConfig.is_default == 1, + AIModelConfig.is_enabled == 1, ) ) return result.scalar_one_or_none() diff --git a/digital-avatar-app/backend/routers/chat.py b/digital-avatar-app/backend/routers/chat.py index e01506d..9e5d863 100644 --- a/digital-avatar-app/backend/routers/chat.py +++ b/digital-avatar-app/backend/routers/chat.py @@ -23,13 +23,10 @@ from services.token_billing import ( reserve_avatar_tokens, settle_reservation, ) +from services.chat_model_config import ChatModelConfig, get_chat_model_config router = APIRouter(tags=["数字分身聊天"]) -CHAT_API_URL = os.getenv("CHAT_API_URL", "https://dashscope.aliyuncs.com/compatible-mode/v1") -CHAT_API_KEY = os.getenv("CHAT_API_KEY", "") -CHAT_MODEL = os.getenv("CHAT_MODEL", "qwen-plus") -CHAT_MAX_OUTPUT_TOKENS = max(128, int(os.getenv("CHAT_MAX_OUTPUT_TOKENS", "1024"))) MAX_MESSAGE_LENGTH = 4000 MAX_HISTORY_MESSAGES = 10 QA_LEXICAL_THRESHOLD = 0.72 @@ -287,22 +284,25 @@ def _search_knowledge(db: Session, avatar_id: str, question: str, top_k: int = 5 return results -def _call_qwen(messages: list[dict], temperature: float) -> dict: - if not CHAT_API_KEY: +def _call_qwen( + messages: list[dict], temperature: float, model_config: ChatModelConfig | None = None +) -> dict: + model_config = model_config or get_chat_model_config() + if not model_config.api_key: raise RuntimeError("Qwen 模型服务未配置 CHAT_API_KEY") - url = f"{CHAT_API_URL.rstrip('/')}/chat/completions" + url = f"{model_config.api_base_url}/chat/completions" payload = { - "model": CHAT_MODEL, + "model": model_config.model, "messages": messages, "temperature": temperature, - "max_tokens": CHAT_MAX_OUTPUT_TOKENS, + "max_tokens": model_config.max_tokens, } try: response = httpx.post( url, - headers={"Authorization": f"Bearer {CHAT_API_KEY}"}, + headers={"Authorization": f"Bearer {model_config.api_key}"}, json=payload, - timeout=30, + timeout=model_config.timeout_seconds, ) response.raise_for_status() data = response.json() @@ -314,21 +314,30 @@ def _call_qwen(messages: list[dict], temperature: float) -> dict: return {"answer": answer.strip(), "usage": data.get("usage") or {}} -def _iter_qwen_stream(messages: list[dict], temperature: float): +def _iter_qwen_stream( + messages: list[dict], temperature: float, model_config: ChatModelConfig | None = None +): """将 OpenAI 兼容接口的 SSE 分片原样转为文本增量。""" - if not CHAT_API_KEY: + model_config = model_config or get_chat_model_config() + if not model_config.api_key: raise RuntimeError("模型服务未配置") - url = f"{CHAT_API_URL.rstrip('/')}/chat/completions" + url = f"{model_config.api_base_url}/chat/completions" payload = { - "model": CHAT_MODEL, + "model": model_config.model, "messages": messages, "temperature": temperature, - "max_tokens": CHAT_MAX_OUTPUT_TOKENS, + "max_tokens": model_config.max_tokens, "stream": True, "stream_options": {"include_usage": True}, } try: - with httpx.stream("POST", url, headers={"Authorization": f"Bearer {CHAT_API_KEY}"}, json=payload, timeout=45) as response: + with httpx.stream( + "POST", + url, + headers={"Authorization": f"Bearer {model_config.api_key}"}, + json=payload, + timeout=max(45, model_config.timeout_seconds), + ) as response: response.raise_for_status() for raw_line in response.iter_lines(): line = raw_line.decode() if isinstance(raw_line, bytes) else raw_line @@ -387,16 +396,21 @@ def _resolve_reply( if model_client is not None: answer = model_client(messages=messages, temperature=temperature) else: + model_config = get_chat_model_config() reservation = reserve_avatar_tokens( db, avatar, usage_source, - CHAT_MODEL, + model_config.model, messages, - CHAT_MAX_OUTPUT_TOKENS, + model_config.max_tokens, ) try: - model_result = _call_qwen(messages=messages, temperature=temperature) + model_result = _call_qwen( + messages=messages, + temperature=temperature, + model_config=model_config, + ) answer = model_result["answer"] token_usage = settle_reservation( db, @@ -436,15 +450,16 @@ def _stream_reply( config = _config(avatar) temperature = min(0.45 if references else 0.25, 0.2 + config["creativity"] / 100 * 0.6) messages = _build_prompt(avatar, history, question, references) + model_config = get_chat_model_config() reservation = reserve_avatar_tokens( db, avatar, usage_source, - CHAT_MODEL, + model_config.model, messages, - CHAT_MAX_OUTPUT_TOKENS, + model_config.max_tokens, ) - chunks = _iter_qwen_stream(messages, temperature) + chunks = _iter_qwen_stream(messages, temperature, model_config) if matched: messages, reservation = [], None if public: diff --git a/digital-avatar-app/backend/services/chat_model_config.py b/digital-avatar-app/backend/services/chat_model_config.py new file mode 100644 index 0000000..a2c71b1 --- /dev/null +++ b/digital-avatar-app/backend/services/chat_model_config.py @@ -0,0 +1,93 @@ +import logging +import os +import threading +import time +from dataclasses import dataclass + +import httpx + +logger = logging.getLogger(__name__) + + +@dataclass(frozen=True) +class ChatModelConfig: + api_base_url: str + api_key: str + model: str + max_tokens: int + timeout_seconds: float + source: str + + +_cache_lock = threading.Lock() +_cached_config: ChatModelConfig | None = None +_cache_expires_at = 0.0 + + +def _environment_config() -> ChatModelConfig: + return ChatModelConfig( + api_base_url=os.getenv( + "CHAT_API_URL", "https://dashscope.aliyuncs.com/compatible-mode/v1" + ).rstrip("/"), + api_key=os.getenv("CHAT_API_KEY", ""), + model=os.getenv("CHAT_MODEL", "qwen-plus"), + max_tokens=max(128, int(os.getenv("CHAT_MAX_OUTPUT_TOKENS", "1024"))), + timeout_seconds=max(5.0, float(os.getenv("CHAT_TIMEOUT_SECONDS", "30"))), + source="environment", + ) + + +def _fetch_runtime_config() -> ChatModelConfig | None: + url = os.getenv("CHAT_MODEL_CONFIG_URL", "").strip() + token = os.getenv("AVATAR_MODEL_CONFIG_TOKEN", "").strip() + if not url or not token: + return None + response = httpx.get( + url, + headers={"X-Avatar-Config-Token": token}, + timeout=max(2.0, float(os.getenv("CHAT_MODEL_CONFIG_TIMEOUT_SECONDS", "5"))), + ) + response.raise_for_status() + payload = response.json().get("data") or {} + api_base_url = str(payload.get("api_base_url") or "").rstrip("/") + api_key = str(payload.get("api_key") or "") + model = str(payload.get("model") or "") + if not api_base_url or not api_key or not model: + raise ValueError("数字分身专用模型配置不完整") + return ChatModelConfig( + api_base_url=api_base_url, + api_key=api_key, + model=model, + max_tokens=max(128, int(payload.get("max_tokens") or 1024)), + timeout_seconds=max(5.0, float(payload.get("timeout_seconds") or 30)), + source="admin", + ) + + +def get_chat_model_config(*, force_refresh: bool = False) -> ChatModelConfig: + global _cached_config, _cache_expires_at + + now = time.monotonic() + if not force_refresh and _cached_config is not None and now < _cache_expires_at: + return _cached_config + + with _cache_lock: + now = time.monotonic() + if not force_refresh and _cached_config is not None and now < _cache_expires_at: + return _cached_config + try: + config = _fetch_runtime_config() or _environment_config() + except (httpx.HTTPError, ValueError, TypeError) as exc: + logger.warning("读取数字分身专用模型配置失败,暂时使用环境变量配置: %s", exc) + config = _environment_config() + _cached_config = config + ttl = max(5, int(os.getenv("CHAT_MODEL_CONFIG_CACHE_SECONDS", "60"))) + _cache_expires_at = now + ttl + return config + + +def clear_chat_model_config_cache() -> None: + global _cached_config, _cache_expires_at + with _cache_lock: + _cached_config = None + _cache_expires_at = 0.0 diff --git a/digital-avatar-app/backend/tests/test_chat_model_config.py b/digital-avatar-app/backend/tests/test_chat_model_config.py new file mode 100644 index 0000000..43bd77b --- /dev/null +++ b/digital-avatar-app/backend/tests/test_chat_model_config.py @@ -0,0 +1,79 @@ +from unittest.mock import Mock, patch + +import httpx + +from services.chat_model_config import ( + clear_chat_model_config_cache, + get_chat_model_config, +) + + +def setup_function(): + clear_chat_model_config_cache() + + +def teardown_function(): + clear_chat_model_config_cache() + + +def test_admin_runtime_config_takes_priority(monkeypatch): + monkeypatch.setenv("CHAT_MODEL_CONFIG_URL", "http://config.test/runtime") + monkeypatch.setenv("AVATAR_MODEL_CONFIG_TOKEN", "shared-secret") + response = Mock() + response.raise_for_status.return_value = None + response.json.return_value = { + "data": { + "api_base_url": "https://model.test/v1/", + "api_key": "runtime-key", + "model": "avatar-model", + "max_tokens": 2048, + "timeout_seconds": 42, + } + } + + with patch("services.chat_model_config.httpx.get", return_value=response) as request: + config = get_chat_model_config() + + assert config.source == "admin" + assert config.api_base_url == "https://model.test/v1" + assert config.model == "avatar-model" + assert config.max_tokens == 2048 + request.assert_called_once_with( + "http://config.test/runtime", + headers={"X-Avatar-Config-Token": "shared-secret"}, + timeout=5.0, + ) + + +def test_runtime_failure_falls_back_to_environment(monkeypatch): + monkeypatch.setenv("CHAT_MODEL_CONFIG_URL", "http://config.test/runtime") + monkeypatch.setenv("AVATAR_MODEL_CONFIG_TOKEN", "shared-secret") + monkeypatch.setenv("CHAT_API_URL", "https://fallback.test/v1/") + monkeypatch.setenv("CHAT_API_KEY", "fallback-key") + monkeypatch.setenv("CHAT_MODEL", "fallback-model") + monkeypatch.setenv("CHAT_MAX_OUTPUT_TOKENS", "1536") + + request = httpx.Request("GET", "http://config.test/runtime") + with patch( + "services.chat_model_config.httpx.get", + side_effect=httpx.ConnectError("offline", request=request), + ): + config = get_chat_model_config() + + assert config.source == "environment" + assert config.api_base_url == "https://fallback.test/v1" + assert config.api_key == "fallback-key" + assert config.model == "fallback-model" + assert config.max_tokens == 1536 + + +def test_runtime_config_is_cached(monkeypatch): + monkeypatch.setenv("CHAT_MODEL_CONFIG_URL", "") + monkeypatch.setenv("CHAT_MODEL", "first-model") + first = get_chat_model_config() + monkeypatch.setenv("CHAT_MODEL", "second-model") + + second = get_chat_model_config() + + assert first is second + assert second.model == "first-model" diff --git a/digital-avatar-app/docker-compose.yml b/digital-avatar-app/docker-compose.yml index ee90f7b..4a9530f 100644 --- a/digital-avatar-app/docker-compose.yml +++ b/digital-avatar-app/docker-compose.yml @@ -10,6 +10,9 @@ services: environment: DATABASE_URL: sqlite:////data/avatar.db UPLOAD_DIR: /data/uploads + CHAT_MODEL_CONFIG_URL: http://host.docker.internal:8000/api/ai-models/runtime/digital-avatar + extra_hosts: + - "host.docker.internal:host-gateway" volumes: - avatar-data:/data expose: diff --git a/docker-compose.yml b/docker-compose.yml index c9908da..b684dbb 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -17,6 +17,7 @@ services: - REDIS_PORT=6379 - SECRET_KEY=your-secret-key-change-in-production - AES_KEY=your-aes-key-32-chars-change-now! + - AVATAR_MODEL_CONFIG_TOKEN=${AVATAR_MODEL_CONFIG_TOKEN:-} - TZ=Asia/Shanghai - AVATAR_DB_PATH=/app/avatar.db volumes: diff --git a/docker/mysql/init.sql b/docker/mysql/init.sql index 0067ddb..2ebda21 100644 --- a/docker/mysql/init.sql +++ b/docker/mysql/init.sql @@ -94,6 +94,7 @@ CREATE TABLE IF NOT EXISTS `ai_model_configs` ( `id` bigint NOT NULL AUTO_INCREMENT, `model_name` varchar(64) NOT NULL COMMENT '模型名称', `provider` varchar(32) NOT NULL COMMENT 'openai/zhipu/wenxin/qianwen/local', + `usage_scope` varchar(16) NOT NULL DEFAULT 'general' COMMENT '用途:general/digital_avatar', `api_base_url` varchar(256) DEFAULT NULL COMMENT 'API地址', `api_key_enc` varchar(512) DEFAULT NULL COMMENT '加密API Key', `model_version` varchar(64) DEFAULT NULL COMMENT '模型版本', diff --git a/frontend/src/views/AIModels.vue b/frontend/src/views/AIModels.vue index 4e80e11..b1b63e0 100644 --- a/frontend/src/views/AIModels.vue +++ b/frontend/src/views/AIModels.vue @@ -15,6 +15,9 @@ {{ m.model_name }}