From e720baa21e8a08971185283b44e0e2cd7ba0f6f9 Mon Sep 17 00:00:00 2001 From: stefanfeng Date: Fri, 21 Aug 2026 13:21:05 +0800 Subject: [PATCH 01/10] fix(avatar): batch knowledge embedding requests --- digital-avatar-app/backend/embeddings.py | 42 +++++++++++------- .../backend/tests/test_embeddings.py | 44 +++++++++++++++++++ 2 files changed, 70 insertions(+), 16 deletions(-) diff --git a/digital-avatar-app/backend/embeddings.py b/digital-avatar-app/backend/embeddings.py index 6c54232..0e20019 100644 --- a/digital-avatar-app/backend/embeddings.py +++ b/digital-avatar-app/backend/embeddings.py @@ -51,22 +51,32 @@ def embed(texts): if api_url: api_key = os.getenv("EMBEDDING_API_KEY", "") model = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small") - payload = json.dumps({"input": texts, "model": model}).encode("utf-8") - req = urllib.request.Request( - api_url, - data=payload, - headers={ - "Content-Type": "application/json", - "Authorization": f"Bearer {api_key}" if api_key else "", - }, - method="POST", - ) - with urllib.request.urlopen(req, timeout=30) as resp: - data = json.loads(resp.read().decode("utf-8")) - items = data["data"] - if items and "index" in items[0]: - items = sorted(items, key=lambda x: x["index"]) - return [item["embedding"] for item in items] + try: + batch_size = max(1, int(os.getenv("EMBEDDING_BATCH_SIZE", "10"))) + except ValueError: + batch_size = 10 + embeddings = [] + for start in range(0, len(texts), batch_size): + batch = texts[start:start + batch_size] + payload = json.dumps({"input": batch, "model": model}).encode("utf-8") + req = urllib.request.Request( + api_url, + data=payload, + headers={ + "Content-Type": "application/json", + "Authorization": f"Bearer {api_key}" if api_key else "", + }, + method="POST", + ) + with urllib.request.urlopen(req, timeout=30) as resp: + data = json.loads(resp.read().decode("utf-8")) + items = data["data"] + if items and "index" in items[0]: + items = sorted(items, key=lambda x: x["index"]) + if len(items) != len(batch): + raise ValueError("embedding response count does not match request") + embeddings.extend(item["embedding"] for item in items) + return embeddings return _hash_embedding(texts) diff --git a/digital-avatar-app/backend/tests/test_embeddings.py b/digital-avatar-app/backend/tests/test_embeddings.py index 7520631..eafb1d8 100644 --- a/digital-avatar-app/backend/tests/test_embeddings.py +++ b/digital-avatar-app/backend/tests/test_embeddings.py @@ -1,10 +1,26 @@ +import json import os import tempfile import unittest +from unittest.mock import patch import embeddings +class FakeResponse: + def __init__(self, payload): + self.payload = payload + + def __enter__(self): + return self + + def __exit__(self, *_): + return None + + def read(self): + return json.dumps(self.payload).encode("utf-8") + + class TextExtractionTests(unittest.TestCase): def write_text(self, suffix, content): handle = tempfile.NamedTemporaryFile(suffix=suffix, delete=False) @@ -28,5 +44,33 @@ class TextExtractionTests(unittest.TestCase): embeddings.extract_text(path, ".csv") +class RemoteEmbeddingTests(unittest.TestCase): + def test_large_input_is_split_into_provider_safe_batches(self): + texts = [f"chunk-{index}" for index in range(14)] + batch_sizes = [] + + def fake_urlopen(request, timeout): + self.assertEqual(timeout, 30) + payload = json.loads(request.data.decode("utf-8")) + batch_sizes.append(len(payload["input"])) + return FakeResponse({ + "data": [ + {"index": index, "embedding": [float(text.split("-")[1])]} + for index, text in enumerate(payload["input"]) + ] + }) + + with patch.dict(os.environ, { + "EMBEDDING_API_URL": "https://embedding.example/v1/embeddings", + "EMBEDDING_API_KEY": "test-key", + "EMBEDDING_MODEL": "text-embedding-v4", + "EMBEDDING_BATCH_SIZE": "10", + }), patch("embeddings.urllib.request.urlopen", side_effect=fake_urlopen): + result = embeddings.embed(texts) + + self.assertEqual(batch_sizes, [10, 4]) + self.assertEqual(result, [[float(index)] for index in range(14)]) + + if __name__ == "__main__": unittest.main() -- 2.54.0 From 672019830d409b5ab6517a52d5ab91dd75c826d6 Mon Sep 17 00:00:00 2001 From: stefanfeng Date: Fri, 21 Aug 2026 16:21:00 +0800 Subject: [PATCH 02/10] fix(avatar): prevent takeover replies blocking across chats --- digital-avatar-app/backend/main.py | 18 ++++++- .../backend/services/takeover_service.py | 50 +++++++++++++------ .../backend/tests/test_takeover_scheduler.py | 23 +++++---- .../backend/tests/test_takeover_service.py | 34 +++++++++++++ 4 files changed, 98 insertions(+), 27 deletions(-) diff --git a/digital-avatar-app/backend/main.py b/digital-avatar-app/backend/main.py index c34ebac..dea64a5 100644 --- a/digital-avatar-app/backend/main.py +++ b/digital-avatar-app/backend/main.py @@ -143,14 +143,28 @@ def on_startup(): poll_interval = max(0.5, float(os.getenv("BOXIM_POLL_INTERVAL_SECONDS", "1"))) takeover_scheduler = AsyncIOScheduler() takeover_scheduler.add_job( - takeover_service.poll_and_process_messages, + takeover_service.poll_messages, trigger=IntervalTrigger(seconds=poll_interval), id="takeover_message_poll", max_instances=1, coalesce=True, ) + process_interval = max( + 0.25, float(os.getenv("TAKEOVER_PROCESS_INTERVAL_SECONDS", "0.5")) + ) + takeover_scheduler.add_job( + takeover_service.process_reply_tasks, + trigger=IntervalTrigger(seconds=process_interval), + id="takeover_reply_process", + max_instances=1, + coalesce=True, + ) takeover_scheduler.start() - logger.info("BOXIM takeover scheduler started (interval=%ss)", poll_interval) + logger.info( + "BOXIM takeover scheduler started (poll=%ss, process=%ss)", + poll_interval, + process_interval, + ) except Exception as e: stop_takeover_scheduler() logger.warning(f"Failed to initialize takeover scheduler, app will continue without it: {e}") diff --git a/digital-avatar-app/backend/services/takeover_service.py b/digital-avatar-app/backend/services/takeover_service.py index d803345..43e8ad0 100644 --- a/digital-avatar-app/backend/services/takeover_service.py +++ b/digital-avatar-app/backend/services/takeover_service.py @@ -86,24 +86,34 @@ class TakeoverService: self.reply_delay_seconds = reply_delay_seconds self.now = now self._sessions: dict[str, dict] = {} - self._run_lock = asyncio.Lock() + self._poll_lock = asyncio.Lock() + self._process_lock = asyncio.Lock() async def poll_and_process_messages(self): - """Run one complete cycle; polling always happens before reply dispatch.""" - if self._run_lock.locked(): + """Run one complete cycle for callers that do not use the split scheduler.""" + await self.poll_messages() + await self.process_reply_tasks() + + async def poll_messages(self): + """Fetch BOXIM events without blocking reply generation and dispatch.""" + if self._poll_lock.locked(): return - async with self._run_lock: + async with self._poll_lock: self._recover_stuck_tasks() avatar_ids = self._enabled_avatar_ids() self._cancel_disabled_tasks(set(avatar_ids)) for avatar_id in avatar_ids: await self._sync_avatar(avatar_id) - generated = await self._prepare_replies() - if generated: - # Catch a human reply sent while the model was preparing its answer. - for avatar_id in avatar_ids: - await self._sync_avatar(avatar_id) + async def process_reply_tasks(self): + """Generate and send replies independently from BOXIM's long poll.""" + if self._process_lock.locked(): + return + async with self._process_lock: + self._recover_stuck_tasks() + avatar_ids = set(self._enabled_avatar_ids()) + self._cancel_disabled_tasks(avatar_ids) + await self._prepare_replies() await self._dispatch_ready_replies() def _enabled_avatar_ids(self) -> list[str]: @@ -454,11 +464,19 @@ class TakeoverService: finally: db.close() - generated = 0 - for task_id in task_ids: - if await asyncio.to_thread(self._generate_reply, task_id): - generated += 1 - return generated + if not task_ids: + return 0 + + # Each conversation owns its task, so unrelated contacts can generate in + # parallel instead of one slow model response delaying every other peer. + semaphore = asyncio.Semaphore(4) + + async def generate(task_id: str) -> bool: + async with semaphore: + return await asyncio.to_thread(self._generate_reply, task_id) + + results = await asyncio.gather(*(generate(task_id) for task_id in task_ids)) + return sum(bool(result) for result in results) def _generate_reply(self, task_id: str) -> bool: db = self.session_factory() @@ -548,8 +566,8 @@ class TakeoverService: finally: db.close() - for task_id in task_ids: - await self._send_task(task_id) + if task_ids: + await asyncio.gather(*(self._send_task(task_id) for task_id in task_ids)) async def _send_task(self, task_id: str) -> bool: db = self.session_factory() diff --git a/digital-avatar-app/backend/tests/test_takeover_scheduler.py b/digital-avatar-app/backend/tests/test_takeover_scheduler.py index 111c119..fff0bce 100644 --- a/digital-avatar-app/backend/tests/test_takeover_scheduler.py +++ b/digital-avatar-app/backend/tests/test_takeover_scheduler.py @@ -25,7 +25,8 @@ def test_scheduler_uses_boxim_and_restart_safe_service( boxim = MagicMock() mock_boxim_class.return_value = boxim takeover = MagicMock() - takeover.poll_and_process_messages = AsyncMock() + takeover.poll_messages = AsyncMock() + takeover.process_reply_tasks = AsyncMock() mock_takeover_class.return_value = takeover environment = { @@ -46,14 +47,18 @@ def test_scheduler_uses_boxim_and_restart_safe_service( assert config["BOXIM_API_BASE_URL"] == "https://im.example/api" mock_takeover_class.assert_called_once_with(main.SessionLocal, boxim) - scheduler.add_job.assert_called_once() - scheduled_callable = scheduler.add_job.call_args.args[0] - job_options = scheduler.add_job.call_args.kwargs - assert scheduled_callable is takeover.poll_and_process_messages - assert job_options["id"] == "takeover_message_poll" - assert job_options["trigger"].interval.total_seconds() == 1 - assert job_options["max_instances"] == 1 - assert job_options["coalesce"] is True + assert scheduler.add_job.call_count == 2 + poll_call, process_call = scheduler.add_job.call_args_list + assert poll_call.args[0] is takeover.poll_messages + assert poll_call.kwargs["id"] == "takeover_message_poll" + assert poll_call.kwargs["trigger"].interval.total_seconds() == 1 + assert poll_call.kwargs["max_instances"] == 1 + assert poll_call.kwargs["coalesce"] is True + assert process_call.args[0] is takeover.process_reply_tasks + assert process_call.kwargs["id"] == "takeover_reply_process" + assert process_call.kwargs["trigger"].interval.total_seconds() == 0.5 + assert process_call.kwargs["max_instances"] == 1 + assert process_call.kwargs["coalesce"] is True scheduler.start.assert_called_once_with() main.takeover_scheduler = None diff --git a/digital-avatar-app/backend/tests/test_takeover_service.py b/digital-avatar-app/backend/tests/test_takeover_service.py index 2786a25..5a6eed6 100644 --- a/digital-avatar-app/backend/tests/test_takeover_service.py +++ b/digital-avatar-app/backend/tests/test_takeover_service.py @@ -1,6 +1,7 @@ """End-to-end service tests for BOXIM takeover timing and human priority.""" from datetime import datetime, timedelta, timezone +from threading import Barrier from unittest.mock import AsyncMock, patch import pytest @@ -143,6 +144,39 @@ async def test_incoming_message_is_prepared_then_sent_at_three_seconds(service_c db.close() +@pytest.mark.asyncio +async def test_different_contacts_generate_without_blocking_each_other(service_context): + session_factory, service, boxim, clock = service_context + await service.poll_and_process_messages() + boxim.messages.extend( + [ + {"id": 13, "localId": 31, "sendId": 200, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "联系人甲"}, + {"id": 14, "localId": 32, "sendId": 300, "recvId": 100, "sendTime": clock.millis(), "type": 0, "content": "联系人乙"}, + ] + ) + both_generating = Barrier(2, timeout=2) + + def resolve(_db, _avatar, prompt, _history): + both_generating.wait() + return {"answer": f"回复{prompt[-1]}"} + + with patch("routers.chat._resolve_reply", side_effect=resolve): + await service.poll_and_process_messages() + + clock.advance(3) + await service.process_reply_tasks() + assert {(item["peerId"], item["content"]) for item in boxim.sent} == { + ("200", "回复甲"), + ("300", "回复乙"), + } + + db = session_factory() + try: + assert {task.status for task in db.query(TakeoverReplyTask).all()} == {"sent"} + finally: + db.close() + + @pytest.mark.asyncio async def test_read_receipt_failure_does_not_advance_cursor(service_context): session_factory, service, boxim, clock = service_context -- 2.54.0 From 7a0199e685f54edecca4e1bcaf1d339e10599870 Mon Sep 17 00:00:00 2001 From: stefanfeng Date: Tue, 25 Aug 2026 11:44:27 +0800 Subject: [PATCH 03/10] feat(avatar): improve multi-avatar management --- digital-avatar-app/src/App.vue | 109 +---- digital-avatar-app/src/router/index.ts | 24 + .../src/utils/avatar-page-data.d.ts | 5 + .../src/utils/avatar-page-data.js | 5 + .../src/views/AuthorizationManage.vue | 9 +- digital-avatar-app/src/views/AvatarManage.vue | 452 +++--------------- .../src/views/KnowledgeManage.vue | 26 +- digital-avatar-app/src/views/QaPairEdit.vue | 12 +- 8 files changed, 125 insertions(+), 517 deletions(-) diff --git a/digital-avatar-app/src/App.vue b/digital-avatar-app/src/App.vue index c54ecb2..c7c6915 100644 --- a/digital-avatar-app/src/App.vue +++ b/digital-avatar-app/src/App.vue @@ -1,73 +1,10 @@ - - diff --git a/digital-avatar-app/src/router/index.ts b/digital-avatar-app/src/router/index.ts index 90d4377..1edac7d 100644 --- a/digital-avatar-app/src/router/index.ts +++ b/digital-avatar-app/src/router/index.ts @@ -45,6 +45,12 @@ const routes: RouteRecordRaw[] = [ component: () => import('@/views/AuthorizationManage.vue'), meta: { title: '授权管理', requiresAuth: true } }, + { + path: '/avatar/:avatarId/authorization', + name: 'AvatarAuthorizationManage', + component: () => import('@/views/AuthorizationManage.vue'), + meta: { title: '授权管理', requiresAuth: true } + }, { path: '/token/charge', name: 'TokenCharge', @@ -81,6 +87,12 @@ const routes: RouteRecordRaw[] = [ component: () => import('@/views/KnowledgeManage.vue'), meta: { title: '知识库管理', requiresAuth: true } }, + { + path: '/avatar/:avatarId/knowledge', + name: 'AvatarKnowledgeManage', + component: () => import('@/views/KnowledgeManage.vue'), + meta: { title: '知识库管理', requiresAuth: true } + }, { path: '/knowledge/qa/create', name: 'QaPairCreate', @@ -93,6 +105,18 @@ const routes: RouteRecordRaw[] = [ component: () => import('@/views/QaPairEdit.vue'), meta: { title: '编辑问答对', requiresAuth: true } }, + { + path: '/avatar/:avatarId/knowledge/qa/create', + name: 'AvatarQaPairCreate', + component: () => import('@/views/QaPairEdit.vue'), + meta: { title: '添加问答对', requiresAuth: true } + }, + { + path: '/avatar/:avatarId/knowledge/qa/:qaId/edit', + name: 'AvatarQaPairEdit', + component: () => import('@/views/QaPairEdit.vue'), + meta: { title: '编辑问答对', requiresAuth: true } + }, { path: '/login/sms', name: 'SmsLogin', diff --git a/digital-avatar-app/src/utils/avatar-page-data.d.ts b/digital-avatar-app/src/utils/avatar-page-data.d.ts index 883f30c..32bdb1f 100644 --- a/digital-avatar-app/src/utils/avatar-page-data.d.ts +++ b/digital-avatar-app/src/utils/avatar-page-data.d.ts @@ -41,5 +41,10 @@ export function pickAvatarId( currentAvatarId: string | null | undefined, avatars?: AvatarPageRecord[] ): string | null +export function pickScopedAvatarId( + routeAvatarId: string | string[] | null | undefined, + currentAvatarId: string | null | undefined, + avatars?: AvatarPageRecord[] +): string | null export function normalizeAvatarEditForm(avatar?: AvatarPageRecord): AvatarEditForm export function buildAvatarUpdatePayload(form: AvatarEditForm): AvatarUpdatePayload diff --git a/digital-avatar-app/src/utils/avatar-page-data.js b/digital-avatar-app/src/utils/avatar-page-data.js index 57be97f..d96fdb9 100644 --- a/digital-avatar-app/src/utils/avatar-page-data.js +++ b/digital-avatar-app/src/utils/avatar-page-data.js @@ -8,6 +8,11 @@ export function pickAvatarId(currentAvatarId, avatars) { return currentAvatarId || avatars?.[0]?.id || null } +export function pickScopedAvatarId(routeAvatarId, currentAvatarId, avatars) { + const requested = Array.isArray(routeAvatarId) ? routeAvatarId[0] : routeAvatarId + return requested ? String(requested) : pickAvatarId(currentAvatarId, avatars) +} + export function normalizeAvatarEditForm(avatar = {}) { const config = avatar.config || {} return { diff --git a/digital-avatar-app/src/views/AuthorizationManage.vue b/digital-avatar-app/src/views/AuthorizationManage.vue index c5095f0..a6510a7 100644 --- a/digital-avatar-app/src/views/AuthorizationManage.vue +++ b/digital-avatar-app/src/views/AuthorizationManage.vue @@ -109,7 +109,7 @@ @@ -234,7 +248,7 @@ onMounted(async () => { .knowledge-page { min-height: 100vh; background: #F8F9FA; - padding-bottom: 80px; + padding-bottom: calc(28px + env(safe-area-inset-bottom)); overflow-x: hidden; } diff --git a/digital-avatar-app/src/views/QaPairEdit.vue b/digital-avatar-app/src/views/QaPairEdit.vue index f9252d3..41161dc 100644 --- a/digital-avatar-app/src/views/QaPairEdit.vue +++ b/digital-avatar-app/src/views/QaPairEdit.vue @@ -49,14 +49,14 @@ import { ref, reactive, computed, onMounted } from 'vue' import { useRouter, useRoute } from 'vue-router' import { useAvatarStore } from '@/store/avatar' -import { pickAvatarId, unwrapListData } from '@/utils/avatar-page-data.js' +import { pickScopedAvatarId, unwrapListData } from '@/utils/avatar-page-data.js' import { getQAPairs, createQAPair, updateQAPair } from '@/api' const router = useRouter() const route = useRoute() const store = useAvatarStore() -const avatarId = computed(() => pickAvatarId(store.currentAvatarId, store.avatars)) +const avatarId = computed(() => pickScopedAvatarId(route.params.avatarId, store.currentAvatarId, store.avatars)) const qaId = computed(() => (route.params.qaId as string) || null) const isEdit = computed(() => !!qaId.value) @@ -104,8 +104,11 @@ const save = async () => { } else { await createQAPair(avatarId.value, payload) } - // 保存成功返回知识库管理页 - router.replace('/knowledge') + if (route.params.avatarId) { + router.replace({ name: 'AvatarKnowledgeManage', params: { avatarId: avatarId.value } }) + } else { + router.replace('/knowledge') + } } catch (e: any) { error.value = e?.message || '保存失败' } finally { @@ -117,6 +120,7 @@ onMounted(async () => { if (!store.avatars.length) { await store.loadAvatars() } + if (avatarId.value) store.currentAvatarId = avatarId.value if (isEdit.value) { await loadForEdit() } -- 2.54.0 From 699bbbde57e51ad2cf293bd0671618c2f3a541b1 Mon Sep 17 00:00:00 2001 From: stefanfeng Date: Tue, 25 Aug 2026 13:24:02 +0800 Subject: [PATCH 04/10] feat(avatar): add user token accounting --- digital-avatar-app/backend/database.py | 14 + digital-avatar-app/backend/main.py | 36 ++- digital-avatar-app/backend/models.py | 33 ++- digital-avatar-app/backend/routers/chat.py | 146 ++++++++-- .../backend/routers/huihui_auth.py | 3 + digital-avatar-app/backend/routers/tokens.py | 70 ++++- .../backend/services/takeover_service.py | 2 +- .../backend/services/token_billing.py | 198 ++++++++++++++ digital-avatar-app/backend/tests/conftest.py | 9 + .../backend/tests/test_takeover_service.py | 8 +- .../backend/tests/test_token_billing.py | 252 ++++++++++++++++++ digital-avatar-app/src/api/index.ts | 21 +- digital-avatar-app/src/store/avatar.ts | 22 +- digital-avatar-app/src/views/AvatarManage.vue | 18 ++ digital-avatar-app/src/views/TokenCharge.vue | 18 +- 15 files changed, 798 insertions(+), 52 deletions(-) create mode 100644 digital-avatar-app/backend/services/token_billing.py create mode 100644 digital-avatar-app/backend/tests/test_token_billing.py diff --git a/digital-avatar-app/backend/database.py b/digital-avatar-app/backend/database.py index 076e2f0..84e7d10 100644 --- a/digital-avatar-app/backend/database.py +++ b/digital-avatar-app/backend/database.py @@ -40,8 +40,14 @@ def init_db(): ("authorizations", "takeover_mode", "VARCHAR DEFAULT 'immediate'"), ("authorizations", "takeover_delay_seconds", "INTEGER DEFAULT 30"), ("avatars", "share_token", "VARCHAR DEFAULT NULL"), + ("token_account", "user_id", "VARCHAR DEFAULT ''"), + ("token_account", "total_granted", "BIGINT DEFAULT 0"), + ("token_account", "total_consumed", "BIGINT DEFAULT 0"), + ("token_account", "created_at", "TIMESTAMP"), + ("token_account", "updated_at", "TIMESTAMP"), ) _normalize_optional_unique_values() + _create_token_indexes() def _try_add_columns(*cols): @@ -58,3 +64,11 @@ def _try_add_columns(*cols): def _normalize_optional_unique_values(): with engine.begin() as conn: conn.exec_driver_sql("UPDATE avatars SET share_token = NULL WHERE share_token = ''") + + +def _create_token_indexes(): + with engine.begin() as conn: + conn.exec_driver_sql( + "CREATE UNIQUE INDEX IF NOT EXISTS ux_token_account_user_id " + "ON token_account(user_id) WHERE user_id <> ''" + ) diff --git a/digital-avatar-app/backend/main.py b/digital-avatar-app/backend/main.py index dea64a5..57beb2a 100644 --- a/digital-avatar-app/backend/main.py +++ b/digital-avatar-app/backend/main.py @@ -8,7 +8,7 @@ from apscheduler.schedulers.asyncio import AsyncIOScheduler from apscheduler.triggers.interval import IntervalTrigger from database import init_db, SessionLocal -from models import Avatar, Authorization, Organization, TokenAccount, TokenPlan +from models import Avatar, Authorization, Organization, TokenAccount, TokenPlan, User from fastapi.staticfiles import StaticFiles import routers.avatars import routers.tokens @@ -19,6 +19,7 @@ import routers.huihui_auth import routers.chat import routers.takeover from responses import ok +from services.token_billing import DEFAULT_TOKEN_GRANT, release_stale_reservations logger = logging.getLogger(__name__) @@ -56,17 +57,29 @@ def health(): def seed(): db = SessionLocal() try: - if db.query(TokenAccount).first() is None: - db.add(TokenAccount(balance=1250)) + plan_specs = [ + {"id": "1", "name": "基础套餐", "amount": 2_000_000, "price": 10, "badge": "", "desc": "2M Token"}, + {"id": "2", "name": "标准套餐", "amount": 20_000_000, "price": 100, "badge": "常用", "desc": "20M Token"}, + {"id": "3", "name": "专业套餐", "amount": 250_000_000, "price": 1000, "badge": "加赠25%", "desc": "250M Token"}, + {"id": "4", "name": "企业套餐", "amount": 2_500_000_000, "price": 10000, "badge": "企业推荐", "desc": "2500M Token"}, + ] + for spec in plan_specs: + plan = db.query(TokenPlan).filter(TokenPlan.id == spec["id"]).first() + if plan is None: + db.add(TokenPlan(**spec)) + else: + for key, value in spec.items(): + setattr(plan, key, value) - if db.query(TokenPlan).count() == 0: - plans = [ - TokenPlan(id="1", name="新手体验", amount=1000, price=9.9, desc="新手体验"), - TokenPlan(id="2", name="热门套餐", amount=5000, price=39.9, badge="热门"), - TokenPlan(id="3", name="超值套餐", amount=12000, price=89.9, badge="超值"), - TokenPlan(id="4", name="企业推荐", amount=30000, price=199, badge="企业推荐", desc="适合高频使用"), - ] - db.add_all(plans) + for user in db.query(User).all(): + account = db.query(TokenAccount).filter(TokenAccount.user_id == user.id).first() + if account is None: + db.add(TokenAccount( + user_id=user.id, + balance=DEFAULT_TOKEN_GRANT, + total_granted=DEFAULT_TOKEN_GRANT, + total_consumed=0, + )) if db.query(Avatar).count() == 0: avatar = Avatar( @@ -105,6 +118,7 @@ def seed(): db.add_all(orgs) db.commit() + release_stale_reservations(db) finally: db.close() diff --git a/digital-avatar-app/backend/models.py b/digital-avatar-app/backend/models.py index 38b1484..b4ae291 100644 --- a/digital-avatar-app/backend/models.py +++ b/digital-avatar-app/backend/models.py @@ -1,6 +1,7 @@ import uuid from sqlalchemy import ( + BigInteger, Boolean, Column, DateTime, @@ -259,14 +260,42 @@ class KnowledgeChunk(Base): class TokenAccount(Base): __tablename__ = "token_account" id = Column(Integer, primary_key=True) - balance = Column(Integer, default=1250) + user_id = Column(String, nullable=False, default="", index=True) + balance = Column(BigInteger, default=1_000_000) + total_granted = Column(BigInteger, default=1_000_000) + total_consumed = Column(BigInteger, default=0) + created_at = Column(DateTime, server_default=func.now()) + updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now()) + + +class TokenUsage(Base): + __tablename__ = "token_usage" + __table_args__ = ( + Index("ix_token_usage_user_created", "user_id", "created_at"), + Index("ix_token_usage_avatar_created", "avatar_id", "created_at"), + ) + + id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex) + user_id = Column(String, nullable=False, index=True) + avatar_id = Column(String, nullable=False, default="", index=True) + source = Column(String, nullable=False, default="chat") + model = Column(String, default="") + status = Column(String, nullable=False, default="reserved") + reserved_tokens = Column(BigInteger, default=0) + prompt_tokens = Column(BigInteger, default=0) + completion_tokens = Column(BigInteger, default=0) + total_tokens = Column(BigInteger, default=0) + balance_after = Column(BigInteger, default=0) + failure_reason = Column(String, default="") + created_at = Column(DateTime, server_default=func.now()) + settled_at = Column(DateTime) class TokenPlan(Base): __tablename__ = "token_plans" id = Column(String, primary_key=True) name = Column(String, default="") - amount = Column(Integer, default=0) + amount = Column(BigInteger, default=0) price = Column(Float, default=0) badge = Column(String, default="") desc = Column(String, default="") diff --git a/digital-avatar-app/backend/routers/chat.py b/digital-avatar-app/backend/routers/chat.py index 74233e2..e01506d 100644 --- a/digital-avatar-app/backend/routers/chat.py +++ b/digital-avatar-app/backend/routers/chat.py @@ -16,12 +16,20 @@ import embeddings from database import get_db from models import Avatar, KnowledgeChunk, KnowledgeDoc, QAPair, User from responses import ok, fail +from services.token_billing import ( + InsufficientTokensError, + estimate_fallback_usage, + release_reservation, + reserve_avatar_tokens, + settle_reservation, +) 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 @@ -279,7 +287,7 @@ def _search_knowledge(db: Session, avatar_id: str, question: str, top_k: int = 5 return results -def _call_qwen(messages: list[dict], temperature: float) -> str: +def _call_qwen(messages: list[dict], temperature: float) -> dict: if not CHAT_API_KEY: raise RuntimeError("Qwen 模型服务未配置 CHAT_API_KEY") url = f"{CHAT_API_URL.rstrip('/')}/chat/completions" @@ -287,6 +295,7 @@ def _call_qwen(messages: list[dict], temperature: float) -> str: "model": CHAT_MODEL, "messages": messages, "temperature": temperature, + "max_tokens": CHAT_MAX_OUTPUT_TOKENS, } try: response = httpx.post( @@ -302,7 +311,7 @@ def _call_qwen(messages: list[dict], temperature: float) -> str: raise RuntimeError("Qwen 模型服务暂时不可用") from exc if not isinstance(answer, str) or not answer.strip(): raise RuntimeError("Qwen 模型没有返回有效回答") - return answer.strip() + return {"answer": answer.strip(), "usage": data.get("usage") or {}} def _iter_qwen_stream(messages: list[dict], temperature: float): @@ -310,7 +319,14 @@ def _iter_qwen_stream(messages: list[dict], temperature: float): if not CHAT_API_KEY: raise RuntimeError("模型服务未配置") url = f"{CHAT_API_URL.rstrip('/')}/chat/completions" - payload = {"model": CHAT_MODEL, "messages": messages, "temperature": temperature, "stream": True} + payload = { + "model": CHAT_MODEL, + "messages": messages, + "temperature": temperature, + "max_tokens": CHAT_MAX_OUTPUT_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: response.raise_for_status() @@ -322,11 +338,15 @@ def _iter_qwen_stream(messages: list[dict], temperature: float): if data == "[DONE]": return try: - delta = json.loads(data).get("choices", [{}])[0].get("delta", {}).get("content") + parsed = json.loads(data) except (ValueError, IndexError, AttributeError): continue + if parsed.get("usage"): + yield {"usage": parsed["usage"]} + choices = parsed.get("choices") or [] + delta = choices[0].get("delta", {}).get("content") if choices else None if delta: - yield delta + yield {"content": delta} except httpx.HTTPError as exc: raise RuntimeError("模型服务暂时不可用") from exc @@ -350,6 +370,7 @@ def _resolve_reply( qa_pairs: list[Any] | None = None, search_fn: Callable[..., list[dict]] | None = None, model_client: Callable[..., str] | None = None, + usage_source: str = "chat", ) -> dict: if qa_pairs is None: qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all() @@ -362,16 +383,49 @@ def _resolve_reply( messages = _build_prompt(avatar, history, question, hits) config = _config(avatar) temperature = min(0.45 if hits else 0.25, 0.2 + config["creativity"] / 100 * 0.6) - model_client = model_client or _call_qwen - answer = model_client(messages=messages, temperature=temperature) - return { + token_usage = None + if model_client is not None: + answer = model_client(messages=messages, temperature=temperature) + else: + reservation = reserve_avatar_tokens( + db, + avatar, + usage_source, + CHAT_MODEL, + messages, + CHAT_MAX_OUTPUT_TOKENS, + ) + try: + model_result = _call_qwen(messages=messages, temperature=temperature) + answer = model_result["answer"] + token_usage = settle_reservation( + db, + reservation, + model_result.get("usage"), + fallback_total=estimate_fallback_usage(messages, answer), + ) + except Exception as exc: + release_reservation(db, reservation, str(exc)) + raise + result = { "answer": answer, "source": "knowledge" if hits else "qwen", "references": hits, } + if token_usage: + result["tokenUsage"] = token_usage + return result -def _stream_reply(db: Session, avatar: Avatar, question: str, history: list[Any], *, public: bool = False): +def _stream_reply( + db: Session, + avatar: Avatar, + question: str, + history: list[Any], + *, + public: bool = False, + usage_source: str = "chat_stream", +): qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all() matched = _match_standard_qa(question, qa_pairs) if matched: @@ -381,18 +435,62 @@ def _stream_reply(db: Session, avatar: Avatar, question: str, history: list[Any] source = "knowledge" if references else "qwen" config = _config(avatar) temperature = min(0.45 if references else 0.25, 0.2 + config["creativity"] / 100 * 0.6) - chunks = _iter_qwen_stream(_build_prompt(avatar, history, question, references), temperature) + messages = _build_prompt(avatar, history, question, references) + reservation = reserve_avatar_tokens( + db, + avatar, + usage_source, + CHAT_MODEL, + messages, + CHAT_MAX_OUTPUT_TOKENS, + ) + chunks = _iter_qwen_stream(messages, temperature) + if matched: + messages, reservation = [], None if public: source, references = "public", [] def generate(): + output_parts = [] + provider_usage = None + settled = False try: yield _sse("meta", {"source": source, "references": references}) - for content in chunks: + for chunk in chunks: + if reservation is None: + content = chunk + else: + provider_usage = chunk.get("usage") or provider_usage + content = chunk.get("content") + if not content: + continue + output_parts.append(content) yield _sse("delta", {"content": content}) - yield _sse("done", {}) + token_usage = None + if reservation is not None: + answer = "".join(output_parts) + token_usage = settle_reservation( + db, + reservation, + provider_usage, + fallback_total=estimate_fallback_usage(messages, answer), + ) + settled = True + yield _sse("done", {} if public else {"tokenUsage": token_usage}) except RuntimeError as exc: yield _sse("error", {"message": str(exc)}) + finally: + if reservation is not None and not settled: + answer = "".join(output_parts) + if answer: + settle_reservation( + db, + reservation, + provider_usage, + fallback_total=estimate_fallback_usage(messages, answer), + ) + else: + release_reservation(db, reservation, "stream_ended_without_output") return StreamingResponse( generate(), @@ -441,11 +539,14 @@ def get_shared_avatar(share_token: str, db: Session = Depends(get_db)): def public_chat(share_token: str, body: ChatIn = Body(...), db: Session = Depends(get_db)): avatar = _require_shared_avatar(db, share_token) try: - result = _resolve_reply(db, avatar, body.message, body.history) + result = _resolve_reply(db, avatar, body.message, body.history, usage_source="public_chat") # 公开访客无需获知知识文件名、检索分数或内部答复来源。 result["references"] = [] result["source"] = "public" + result.pop("tokenUsage", None) return ok(result) + except InsufficientTokensError as exc: + return fail(str(exc), code=402) except RuntimeError as exc: return fail(str(exc), code=502) @@ -455,15 +556,30 @@ def chat(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(N avatar = _require_owned_avatar(db, avatar_id, authorization) try: return ok(_resolve_reply(db, avatar, body.message, body.history)) + except InsufficientTokensError as exc: + return fail(str(exc), code=402) except RuntimeError as exc: return fail(str(exc), code=502) @router.post("/avatar/{avatar_id}/chat/stream") def chat_stream(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)): - return _stream_reply(db, _require_owned_avatar(db, avatar_id, authorization), body.message, body.history) + try: + return _stream_reply(db, _require_owned_avatar(db, avatar_id, authorization), body.message, body.history) + except InsufficientTokensError as exc: + raise HTTPException(status_code=402, detail=str(exc)) from exc @router.post("/public/avatar/{share_token}/chat/stream") def public_chat_stream(share_token: str, body: ChatIn = Body(...), db: Session = Depends(get_db)): - return _stream_reply(db, _require_shared_avatar(db, share_token), body.message, body.history, public=True) + try: + return _stream_reply( + db, + _require_shared_avatar(db, share_token), + body.message, + body.history, + public=True, + usage_source="public_chat_stream", + ) + except InsufficientTokensError as exc: + raise HTTPException(status_code=402, detail=str(exc)) from exc diff --git a/digital-avatar-app/backend/routers/huihui_auth.py b/digital-avatar-app/backend/routers/huihui_auth.py index 25bf880..5c2ba96 100644 --- a/digital-avatar-app/backend/routers/huihui_auth.py +++ b/digital-avatar-app/backend/routers/huihui_auth.py @@ -349,6 +349,9 @@ def _issue_session(db: Session, phone: str, info: dict): db.commit() db.refresh(user) + from services.token_billing import get_or_create_account + get_or_create_account(db, user.id) + return ok({ "token": user.app_token, "user": user.to_dict(), diff --git a/digital-avatar-app/backend/routers/tokens.py b/digital-avatar-app/backend/routers/tokens.py index 18871e8..3a447e9 100644 --- a/digital-avatar-app/backend/routers/tokens.py +++ b/digital-avatar-app/backend/routers/tokens.py @@ -1,37 +1,81 @@ -from fastapi import APIRouter, Depends, Body +from fastapi import APIRouter, Depends, Body, Header, HTTPException +from sqlalchemy import func from sqlalchemy.orm import Session from database import get_db -from models import TokenAccount, TokenPlan +from models import TokenAccount, TokenPlan, TokenUsage, User from responses import ok, fail +from services.token_billing import get_or_create_account router = APIRouter(tags=["Token"]) +def _require_user(authorization: str | None, db: Session) -> User: + if not authorization: + raise HTTPException(status_code=401, detail="未登录") + token = authorization.replace("Bearer ", "", 1).replace("bearer ", "", 1).strip() + user = db.query(User).filter(User.app_token == token).first() + if not user: + raise HTTPException(status_code=401, detail="会话无效或已过期") + return user + + @router.get("/token/balance") -def balance(db: Session = Depends(get_db)): - acc = db.query(TokenAccount).first() - return ok({"balance": acc.balance if acc else 0}) +def balance(authorization: str = Header(None), db: Session = Depends(get_db)): + user = _require_user(authorization, db) + acc = get_or_create_account(db, user.id) + return ok({ + "balance": acc.balance, + "totalGranted": acc.total_granted, + "totalConsumed": acc.total_consumed, + }) @router.get("/token/plans") -def plans(db: Session = Depends(get_db)): +def plans(authorization: str = Header(None), db: Session = Depends(get_db)): + _require_user(authorization, db) items = db.query(TokenPlan).order_by(TokenPlan.price.asc()).all() return ok([p.to_dict() for p in items]) @router.post("/token/charge") -def charge(payload: dict = Body(...), db: Session = Depends(get_db)): +def charge(payload: dict = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)): + user = _require_user(authorization, db) plan_id = payload.get("planId") plan = db.query(TokenPlan).filter(TokenPlan.id == plan_id).first() if not plan: return fail("套餐不存在", 404) - acc = db.query(TokenAccount).first() - if not acc: - acc = TokenAccount(balance=0) - db.add(acc) - db.commit() - db.refresh(acc) + acc = get_or_create_account(db, user.id) acc.balance += plan.amount + acc.total_granted = int(acc.total_granted or 0) + plan.amount db.commit() return ok({"balance": acc.balance, "charged": plan.amount}) + + +@router.get("/token/usage") +def usage(authorization: str = Header(None), db: Session = Depends(get_db)): + user = _require_user(authorization, db) + rows = ( + db.query( + TokenUsage.avatar_id, + TokenUsage.source, + func.sum(TokenUsage.prompt_tokens), + func.sum(TokenUsage.completion_tokens), + func.sum(TokenUsage.total_tokens), + func.count(TokenUsage.id), + ) + .filter(TokenUsage.user_id == user.id, TokenUsage.status == "completed") + .group_by(TokenUsage.avatar_id, TokenUsage.source) + .all() + ) + return ok([ + { + "avatarId": avatar_id, + "source": source, + "promptTokens": int(prompt_tokens or 0), + "completionTokens": int(completion_tokens or 0), + "totalTokens": int(total_tokens or 0), + "requestCount": int(request_count or 0), + } + for avatar_id, source, prompt_tokens, completion_tokens, total_tokens, request_count in rows + ]) diff --git a/digital-avatar-app/backend/services/takeover_service.py b/digital-avatar-app/backend/services/takeover_service.py index 43e8ad0..9a96f9d 100644 --- a/digital-avatar-app/backend/services/takeover_service.py +++ b/digital-avatar-app/backend/services/takeover_service.py @@ -520,7 +520,7 @@ class TakeoverService: from routers.chat import _resolve_reply - result = _resolve_reply(db, avatar, task.prompt, history) + result = _resolve_reply(db, avatar, task.prompt, history, usage_source="takeover") answer = _plain_text_reply(result.get("answer", "")) db.refresh(task) if task.status != "generating": diff --git a/digital-avatar-app/backend/services/token_billing.py b/digital-avatar-app/backend/services/token_billing.py new file mode 100644 index 0000000..ab763a1 --- /dev/null +++ b/digital-avatar-app/backend/services/token_billing.py @@ -0,0 +1,198 @@ +"""User-scoped token accounting for every avatar model request.""" + +import math +from dataclasses import dataclass +from datetime import datetime, timedelta + +from sqlalchemy.exc import IntegrityError +from sqlalchemy.orm import Session + +from models import Avatar, TokenAccount, TokenUsage, User + +DEFAULT_TOKEN_GRANT = 1_000_000 + + +class InsufficientTokensError(RuntimeError): + pass + + +@dataclass(frozen=True) +class TokenReservation: + usage_id: str + user_id: str + reserved_tokens: int + + +def get_or_create_account(db: Session, user_id: str) -> TokenAccount: + account = db.query(TokenAccount).filter(TokenAccount.user_id == user_id).first() + if account: + return account + account = TokenAccount( + user_id=user_id, + balance=DEFAULT_TOKEN_GRANT, + total_granted=DEFAULT_TOKEN_GRANT, + total_consumed=0, + ) + db.add(account) + try: + db.commit() + except IntegrityError: + # A concurrent first request may have created the same user account. + db.rollback() + account = db.query(TokenAccount).filter(TokenAccount.user_id == user_id).first() + if account is None: + raise + db.refresh(account) + return account + + +def avatar_owner_user(db: Session, avatar: Avatar) -> User | None: + owner_id = (avatar.owner_id or "").strip() + if not owner_id: + return None + return db.query(User).filter(User.huihui_user_id == owner_id).first() + + +def estimate_request_tokens(messages: list[dict], max_output_tokens: int) -> int: + # UTF-8 bytes / 2 deliberately overestimates mixed Chinese/English prompts; + # the unused reservation is returned after provider usage is received. + content_bytes = sum( + len(str(item.get("content", "")).encode("utf-8")) + for item in messages + ) + prompt_reserve = max(1, math.ceil(content_bytes / 2) + len(messages) * 6) + return prompt_reserve + max(1, int(max_output_tokens)) + + +def estimate_fallback_usage(messages: list[dict], output: str) -> int: + content_bytes = sum( + len(str(item.get("content", "")).encode("utf-8")) + for item in messages + ) + len((output or "").encode("utf-8")) + return max(1, math.ceil(content_bytes / 3) + len(messages) * 4) + + +def reserve_avatar_tokens( + db: Session, + avatar: Avatar, + source: str, + model: str, + messages: list[dict], + max_output_tokens: int, +) -> TokenReservation: + user = avatar_owner_user(db, avatar) + if not user: + raise InsufficientTokensError("分身尚未关联有效用户,暂时无法使用 Token") + account = get_or_create_account(db, user.id) + reserved = estimate_request_tokens(messages, max_output_tokens) + updated = ( + db.query(TokenAccount) + .filter(TokenAccount.id == account.id, TokenAccount.balance >= reserved) + .update( + {TokenAccount.balance: TokenAccount.balance - reserved}, + synchronize_session=False, + ) + ) + if updated != 1: + db.rollback() + raise InsufficientTokensError("Token 余额不足,请充值后继续") + db.refresh(account) + usage = TokenUsage( + user_id=user.id, + avatar_id=avatar.id, + source=source, + model=model, + status="reserved", + reserved_tokens=reserved, + ) + db.add(usage) + db.flush() + usage.balance_after = account.balance + db.commit() + return TokenReservation(usage.id, user.id, reserved) + + +def settle_reservation( + db: Session, + reservation: TokenReservation, + usage: dict | None, + *, + fallback_total: int, +) -> dict: + record = db.query(TokenUsage).filter(TokenUsage.id == reservation.usage_id).first() + if not record or record.status != "reserved": + return {} + provider_usage = usage or {} + prompt_tokens = max(0, int(provider_usage.get("prompt_tokens") or 0)) + completion_tokens = max(0, int(provider_usage.get("completion_tokens") or 0)) + provider_total = max( + int(provider_usage.get("total_tokens") or 0), + prompt_tokens + completion_tokens, + ) + total_tokens = max(1, provider_total or int(fallback_total or 0)) + updated = ( + db.query(TokenAccount) + .filter(TokenAccount.user_id == reservation.user_id) + .update( + { + TokenAccount.balance: TokenAccount.balance + reservation.reserved_tokens - total_tokens, + TokenAccount.total_consumed: TokenAccount.total_consumed + total_tokens, + }, + synchronize_session=False, + ) + ) + if updated != 1: + raise RuntimeError("Token 账户不存在") + db.expire_all() + account = db.query(TokenAccount).filter(TokenAccount.user_id == reservation.user_id).first() + record.prompt_tokens = prompt_tokens + record.completion_tokens = completion_tokens + record.total_tokens = total_tokens + record.balance_after = account.balance + record.status = "completed" + record.settled_at = datetime.utcnow() + db.commit() + return { + "promptTokens": prompt_tokens, + "completionTokens": completion_tokens, + "totalTokens": total_tokens, + "balance": account.balance, + } + + +def release_reservation(db: Session, reservation: TokenReservation, reason: str = "") -> None: + record = db.query(TokenUsage).filter(TokenUsage.id == reservation.usage_id).first() + if not record or record.status != "reserved": + return + updated = ( + db.query(TokenAccount) + .filter(TokenAccount.user_id == reservation.user_id) + .update( + {TokenAccount.balance: TokenAccount.balance + reservation.reserved_tokens}, + synchronize_session=False, + ) + ) + if updated: + db.expire_all() + account = db.query(TokenAccount).filter(TokenAccount.user_id == reservation.user_id).first() + record = db.query(TokenUsage).filter(TokenUsage.id == reservation.usage_id).first() + record.balance_after = account.balance + record.status = "failed" + record.failure_reason = (reason or "model_request_failed")[:255] + record.settled_at = datetime.utcnow() + db.commit() + + +def release_stale_reservations(db: Session, older_than_minutes: int = 10) -> int: + cutoff = datetime.utcnow() - timedelta(minutes=older_than_minutes) + stale = db.query(TokenUsage).filter( + TokenUsage.status == "reserved", + TokenUsage.created_at < cutoff, + ).all() + for record in stale: + release_reservation( + db, + TokenReservation(record.id, record.user_id, int(record.reserved_tokens or 0)), + "stale_reservation_recovered", + ) + return len(stale) diff --git a/digital-avatar-app/backend/tests/conftest.py b/digital-avatar-app/backend/tests/conftest.py index 9d2068f..53d78a8 100644 --- a/digital-avatar-app/backend/tests/conftest.py +++ b/digital-avatar-app/backend/tests/conftest.py @@ -8,6 +8,8 @@ from models import ( TakeoverCursor, TakeoverMessage, TakeoverReplyTask, + TokenAccount, + TokenUsage, User, ) @@ -107,6 +109,13 @@ def authorization_context(): db.query(Avatar).filter(Avatar.id.in_(avatar_ids)).delete( synchronize_session=False ) + user_ids = [owner.id, other.id] + db.query(TokenUsage).filter(TokenUsage.user_id.in_(user_ids)).delete( + synchronize_session=False + ) + db.query(TokenAccount).filter(TokenAccount.user_id.in_(user_ids)).delete( + synchronize_session=False + ) db.query(User).filter(User.id.in_([owner.id, other.id])).delete( synchronize_session=False ) diff --git a/digital-avatar-app/backend/tests/test_takeover_service.py b/digital-avatar-app/backend/tests/test_takeover_service.py index 5a6eed6..f8e850f 100644 --- a/digital-avatar-app/backend/tests/test_takeover_service.py +++ b/digital-avatar-app/backend/tests/test_takeover_service.py @@ -7,7 +7,6 @@ from unittest.mock import AsyncMock, patch import pytest from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker -from sqlalchemy.pool import StaticPool from database import Base from models import Avatar, TakeoverCursor, TakeoverMessage, TakeoverReplyTask, User @@ -59,11 +58,10 @@ class FakeBoxIM: @pytest.fixture -def service_context(): +def service_context(tmp_path): engine = create_engine( - "sqlite://", + f"sqlite:///{tmp_path / 'takeover.db'}", connect_args={"check_same_thread": False}, - poolclass=StaticPool, ) session_factory = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False) Base.metadata.create_all(engine) @@ -156,7 +154,7 @@ async def test_different_contacts_generate_without_blocking_each_other(service_c ) both_generating = Barrier(2, timeout=2) - def resolve(_db, _avatar, prompt, _history): + def resolve(_db, _avatar, prompt, _history, **_kwargs): both_generating.wait() return {"answer": f"回复{prompt[-1]}"} diff --git a/digital-avatar-app/backend/tests/test_token_billing.py b/digital-avatar-app/backend/tests/test_token_billing.py new file mode 100644 index 0000000..3860a9d --- /dev/null +++ b/digital-avatar-app/backend/tests/test_token_billing.py @@ -0,0 +1,252 @@ +import uuid +from concurrent.futures import ThreadPoolExecutor +from threading import Barrier +from unittest.mock import patch + +import pytest +from fastapi.testclient import TestClient + +from database import SessionLocal +from main import app, seed +from models import Avatar, TokenAccount, TokenPlan, TokenUsage, User +from routers.chat import _resolve_reply, _stream_reply +from services.token_billing import ( + DEFAULT_TOKEN_GRANT, + InsufficientTokensError, + get_or_create_account, + release_reservation, + reserve_avatar_tokens, + settle_reservation, +) + +client = TestClient(app) + + +def test_balance_is_user_scoped_and_defaults_to_one_million(authorization_context): + context = authorization_context + owner = client.get("/api/token/balance", headers=context["owner_headers"]) + other = client.get("/api/token/balance", headers=context["other_headers"]) + + assert owner.status_code == 200 + assert owner.json()["data"] == { + "balance": DEFAULT_TOKEN_GRANT, + "totalGranted": DEFAULT_TOKEN_GRANT, + "totalConsumed": 0, + } + assert other.json()["data"]["balance"] == DEFAULT_TOKEN_GRANT + assert client.get("/api/token/balance").status_code == 401 + + +def test_seed_synchronizes_requested_recharge_plans(): + seed() + db = SessionLocal() + try: + plans = db.query(TokenPlan).order_by(TokenPlan.price.asc()).all() + assert [(plan.price, plan.amount) for plan in plans] == [ + (10, 2_000_000), + (100, 20_000_000), + (1000, 250_000_000), + (10000, 2_500_000_000), + ] + finally: + db.close() + + +def test_multiple_avatars_share_owner_balance_and_usage_is_itemized(): + suffix = uuid.uuid4().hex + db = SessionLocal() + user = User(id=f"token-user-{suffix}", huihui_user_id=f"token-huihui-{suffix}") + first = Avatar(id=f"token-avatar-a-{suffix}", owner_id=user.huihui_user_id, name="甲") + second = Avatar(id=f"token-avatar-b-{suffix}", owner_id=user.huihui_user_id, name="乙") + db.add_all([user, first, second]) + db.commit() + try: + first_reservation = reserve_avatar_tokens(db, first, "chat", "qwen-test", [{"content": "问题一"}], 128) + settle_reservation( + db, + first_reservation, + {"prompt_tokens": 60, "completion_tokens": 40, "total_tokens": 100}, + fallback_total=999, + ) + second_reservation = reserve_avatar_tokens(db, second, "takeover", "qwen-test", [{"content": "问题二"}], 128) + settle_reservation( + db, + second_reservation, + {"prompt_tokens": 120, "completion_tokens": 80, "total_tokens": 200}, + fallback_total=999, + ) + + account = get_or_create_account(db, user.id) + assert account.balance == DEFAULT_TOKEN_GRANT - 300 + assert account.total_consumed == 300 + usages = db.query(TokenUsage).filter(TokenUsage.user_id == user.id).order_by(TokenUsage.total_tokens).all() + assert [(row.avatar_id, row.source, row.total_tokens) for row in usages] == [ + (first.id, "chat", 100), + (second.id, "takeover", 200), + ] + finally: + db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False) + db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False) + db.query(Avatar).filter(Avatar.id.in_([first.id, second.id])).delete(synchronize_session=False) + db.query(User).filter(User.id == user.id).delete(synchronize_session=False) + db.commit() + db.close() + + +def test_concurrent_settlements_do_not_overwrite_each_other(): + suffix = uuid.uuid4().hex + db = SessionLocal() + user = User(id=f"concurrent-user-{suffix}", huihui_user_id=f"concurrent-huihui-{suffix}") + avatar = Avatar(id=f"concurrent-avatar-{suffix}", owner_id=user.huihui_user_id, name="并发测试") + db.add_all([user, avatar]) + db.commit() + first = reserve_avatar_tokens(db, avatar, "takeover", "qwen-test", [{"content": "甲"}], 128) + second = reserve_avatar_tokens(db, avatar, "takeover", "qwen-test", [{"content": "乙"}], 128) + db.close() + barrier = Barrier(2, timeout=3) + + def settle(reservation, total): + thread_db = SessionLocal() + try: + barrier.wait() + settle_reservation( + thread_db, + reservation, + {"prompt_tokens": total - 20, "completion_tokens": 20, "total_tokens": total}, + fallback_total=999, + ) + finally: + thread_db.close() + + with ThreadPoolExecutor(max_workers=2) as pool: + list(pool.map(lambda args: settle(*args), [(first, 100), (second, 200)])) + + db = SessionLocal() + try: + account = get_or_create_account(db, user.id) + assert account.balance == DEFAULT_TOKEN_GRANT - 300 + assert account.total_consumed == 300 + finally: + db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False) + db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False) + db.query(Avatar).filter(Avatar.id == avatar.id).delete(synchronize_session=False) + db.query(User).filter(User.id == user.id).delete(synchronize_session=False) + db.commit() + db.close() + + +def test_failed_model_request_returns_the_full_reservation(): + suffix = uuid.uuid4().hex + db = SessionLocal() + user = User(id=f"refund-user-{suffix}", huihui_user_id=f"refund-huihui-{suffix}") + avatar = Avatar(id=f"refund-avatar-{suffix}", owner_id=user.huihui_user_id, name="退款测试") + db.add_all([user, avatar]) + db.commit() + try: + reservation = reserve_avatar_tokens(db, avatar, "chat", "qwen-test", [{"content": "问题"}], 128) + release_reservation(db, reservation, "provider error") + account = get_or_create_account(db, user.id) + usage = db.query(TokenUsage).filter(TokenUsage.id == reservation.usage_id).one() + assert account.balance == DEFAULT_TOKEN_GRANT + assert account.total_consumed == 0 + assert usage.status == "failed" + finally: + db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False) + db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False) + db.query(Avatar).filter(Avatar.id == avatar.id).delete(synchronize_session=False) + db.query(User).filter(User.id == user.id).delete(synchronize_session=False) + db.commit() + db.close() + + +def test_insufficient_balance_rejects_before_model_usage_is_created(): + suffix = uuid.uuid4().hex + db = SessionLocal() + user = User(id=f"empty-user-{suffix}", huihui_user_id=f"empty-huihui-{suffix}") + avatar = Avatar(id=f"empty-avatar-{suffix}", owner_id=user.huihui_user_id, name="余额不足") + db.add_all([user, avatar]) + db.commit() + try: + account = get_or_create_account(db, user.id) + account.balance = 1 + db.commit() + with pytest.raises(InsufficientTokensError): + reserve_avatar_tokens(db, avatar, "chat", "qwen-test", [{"content": "问题"}], 128) + db.refresh(account) + assert account.balance == 1 + assert db.query(TokenUsage).filter(TokenUsage.user_id == user.id).count() == 0 + finally: + db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False) + db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False) + db.query(Avatar).filter(Avatar.id == avatar.id).delete(synchronize_session=False) + db.query(User).filter(User.id == user.id).delete(synchronize_session=False) + db.commit() + db.close() + + +def test_chat_settles_from_provider_usage_not_fallback_estimate(): + suffix = uuid.uuid4().hex + db = SessionLocal() + user = User(id=f"chat-user-{suffix}", huihui_user_id=f"chat-huihui-{suffix}") + avatar = Avatar(id=f"chat-avatar-{suffix}", owner_id=user.huihui_user_id, name="聊天测试", config={}) + db.add_all([user, avatar]) + db.commit() + try: + with patch( + "routers.chat._call_qwen", + return_value={ + "answer": "测试回答", + "usage": {"prompt_tokens": 80, "completion_tokens": 20, "total_tokens": 100}, + }, + ): + result = _resolve_reply( + db, + avatar, + "测试问题", + [], + qa_pairs=[], + search_fn=lambda *_args: [], + ) + assert result["tokenUsage"]["totalTokens"] == 100 + assert result["tokenUsage"]["balance"] == DEFAULT_TOKEN_GRANT - 100 + finally: + db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False) + db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False) + db.query(Avatar).filter(Avatar.id == avatar.id).delete(synchronize_session=False) + db.query(User).filter(User.id == user.id).delete(synchronize_session=False) + db.commit() + db.close() + + +@pytest.mark.asyncio +async def test_streaming_chat_settles_final_provider_usage(): + suffix = uuid.uuid4().hex + db = SessionLocal() + user = User(id=f"stream-user-{suffix}", huihui_user_id=f"stream-huihui-{suffix}") + avatar = Avatar(id=f"stream-avatar-{suffix}", owner_id=user.huihui_user_id, name="流式测试", config={}) + db.add_all([user, avatar]) + db.commit() + try: + chunks = iter([ + {"content": "流式"}, + {"content": "回答"}, + {"usage": {"prompt_tokens": 90, "completion_tokens": 10, "total_tokens": 100}}, + ]) + with patch("routers.chat._iter_qwen_stream", return_value=chunks): + response = _stream_reply(db, avatar, "测试问题", []) + body = [] + async for chunk in response.body_iterator: + body.append(chunk.decode() if isinstance(chunk, bytes) else chunk) + assert "流式" in "".join(body) + account = get_or_create_account(db, user.id) + usage = db.query(TokenUsage).filter(TokenUsage.user_id == user.id).one() + assert account.balance == DEFAULT_TOKEN_GRANT - 100 + assert usage.source == "chat_stream" + assert usage.total_tokens == 100 + finally: + db.query(TokenUsage).filter(TokenUsage.user_id == user.id).delete(synchronize_session=False) + db.query(TokenAccount).filter(TokenAccount.user_id == user.id).delete(synchronize_session=False) + db.query(Avatar).filter(Avatar.id == avatar.id).delete(synchronize_session=False) + db.query(User).filter(User.id == user.id).delete(synchronize_session=False) + db.commit() + db.close() diff --git a/digital-avatar-app/src/api/index.ts b/digital-avatar-app/src/api/index.ts index 59f0166..55daf23 100644 --- a/digital-avatar-app/src/api/index.ts +++ b/digital-avatar-app/src/api/index.ts @@ -131,9 +131,24 @@ export const deleteAvatar = (id: string) => // ==================== Token 管理 API ==================== +export interface TokenBalance { + balance: number + totalGranted: number + totalConsumed: number +} + +export interface TokenUsageSummary { + avatarId: string + source: string + promptTokens: number + completionTokens: number + totalTokens: number + requestCount: number +} + // 获取 Token 余额 export const getTokenBalance = () => - request.get<{ balance: number }>('/token/balance') + request.get('/token/balance') // 获取充值套餐 export const getRechargePlans = () => @@ -143,6 +158,10 @@ export const getRechargePlans = () => export const chargeToken = (planId: string) => request.post<{ balance: number; charged: number }>('/token/charge', { planId }) +// 按分身和使用场景汇总 Token 消耗 +export const getTokenUsage = () => + request.get('/token/usage') + // ==================== 授权管理 API ==================== export type AvatarPermission = 'friend' | 'chat' | 'publish' | 'browse' | 'interact' | 'takeover' diff --git a/digital-avatar-app/src/store/avatar.ts b/digital-avatar-app/src/store/avatar.ts index 0af7f4f..ed303a4 100644 --- a/digital-avatar-app/src/store/avatar.ts +++ b/digital-avatar-app/src/store/avatar.ts @@ -1,13 +1,15 @@ import { defineStore } from 'pinia' import { ref } from 'vue' -import { getAvatarList, createAvatar as apiCreate, deleteAvatar as apiDelete, getTokenBalance, getUserProfile } from '@/api' +import { getAvatarList, createAvatar as apiCreate, deleteAvatar as apiDelete, getTokenBalance, getTokenUsage, getUserProfile } from '@/api' import { unwrapListData } from '@/utils/avatar-page-data' export const useAvatarStore = defineStore('avatar', () => { // 已创建的分身列表(来自后端) const avatars = ref([]) - // 全局 Token 余额(来自后端) + // 当前用户所有分身共享的 Token 账户 const tokenBalance = ref(0) + const tokenConsumed = ref(0) + const tokenUsageByAvatar = ref>({}) // 当前选中分身 id const currentAvatarId = ref(null) // 会会用户资料(头像/昵称,来自会会接口) @@ -29,11 +31,24 @@ export const useAvatarStore = defineStore('avatar', () => { try { const res = await getTokenBalance() tokenBalance.value = (res as any)?.balance ?? 0 + tokenConsumed.value = (res as any)?.totalConsumed ?? 0 } catch (e) { console.error('加载余额失败', e) } } + const loadTokenUsage = async () => { + try { + const rows = await getTokenUsage() + tokenUsageByAvatar.value = rows.reduce>((result, row) => { + result[row.avatarId] = (result[row.avatarId] || 0) + row.totalTokens + return result + }, {}) + } catch (e) { + console.error('加载 Token 用量失败', e) + } + } + // 拉取会会用户资料(头像/昵称) const loadUserProfile = async () => { // 若已通过 uniapp 壳注入(混合架构),优先保留,不回退到后端 mock @@ -82,10 +97,13 @@ export const useAvatarStore = defineStore('avatar', () => { return { avatars, tokenBalance, + tokenConsumed, + tokenUsageByAvatar, currentAvatarId, userProfile, loadAvatars, loadTokenBalance, + loadTokenUsage, loadUserProfile, setNativeProfile, addAvatar, diff --git a/digital-avatar-app/src/views/AvatarManage.vue b/digital-avatar-app/src/views/AvatarManage.vue index 117e3a2..28a5805 100644 --- a/digital-avatar-app/src/views/AvatarManage.vue +++ b/digital-avatar-app/src/views/AvatarManage.vue @@ -30,6 +30,7 @@
Token 余额 {{ tokenBalance.toLocaleString() }} + 累计使用 {{ tokenConsumed.toLocaleString() }}
@@ -55,6 +56,7 @@

{{ a.displayName || a.name }}

{{ statusText(a.status) }}

{{ a.description || '暂无描述' }}

+ 累计使用 {{ avatarTokenUsage(a.id).toLocaleString() }} Token
@@ -95,7 +97,9 @@ const me = computed(() => userStore.user) // 状态(来自 store / 后端) const tokenBalance = computed(() => avatarStore.tokenBalance) +const tokenConsumed = computed(() => avatarStore.tokenConsumed) const avatars = computed(() => avatarStore.avatars) +const avatarTokenUsage = (id: string) => avatarStore.tokenUsageByAvatar[id] || 0 const shareToast = ref('') @@ -174,6 +178,7 @@ onMounted(() => { userStore.loadFromStorage() avatarStore.loadAvatars() avatarStore.loadTokenBalance() + avatarStore.loadTokenUsage() }) @@ -316,6 +321,12 @@ onMounted(() => { color: #F97316; } +.token-used { + margin-top: 3px; + color: #A0A5B4; + font-size: 11px; +} + .recharge-btn { padding: 8px 16px; background: #F97316; @@ -440,6 +451,13 @@ onMounted(() => { white-space: nowrap; } +.avatar-token-usage { + display: inline-block; + margin-top: 5px; + color: #A0A5B4; + font-size: 10px; +} + .avatar-status { display: inline-flex; align-items: center; diff --git a/digital-avatar-app/src/views/TokenCharge.vue b/digital-avatar-app/src/views/TokenCharge.vue index f621b5c..5447e65 100644 --- a/digital-avatar-app/src/views/TokenCharge.vue +++ b/digital-avatar-app/src/views/TokenCharge.vue @@ -13,6 +13,7 @@ 当前余额 {{ currentBalance.toLocaleString() }} Token + 累计使用 {{ totalConsumed.toLocaleString() }} Token
@@ -28,7 +29,7 @@ @click="selectedPlan = plan" >
{{ plan.badge }}
-
{{ plan.amount.toLocaleString() }}
+
{{ formatTokenAmount(plan.amount) }}
Token
¥{{ plan.price }}
{{ plan.desc }}
@@ -83,7 +84,8 @@ import { getTokenBalance, getRechargePlans, chargeToken } from '@/api' const router = useRouter() // 当前余额 -const currentBalance = ref(1250) +const currentBalance = ref(0) +const totalConsumed = ref(0) // 充值套餐 const plans = ref { try { const b: any = await getTokenBalance() currentBalance.value = b?.balance ?? 0 + totalConsumed.value = b?.totalConsumed ?? 0 } catch (e) { console.error('加载余额失败', e) } @@ -118,6 +121,10 @@ const loadData = async () => { // 执行充值(写入后端) const charging = ref(false) +const formatTokenAmount = (amount: number) => { + if (amount >= 1_000_000 && amount % 1_000_000 === 0) return `${amount / 1_000_000}M` + return amount.toLocaleString() +} const doCharge = async () => { if (!selectedPlan.value || charging.value) return charging.value = true @@ -187,6 +194,7 @@ onMounted(() => { .balance-card { display: flex; + flex-wrap: wrap; align-items: baseline; gap: 8px; padding: 20px; @@ -211,6 +219,12 @@ onMounted(() => { opacity: 0.9; } +.balance-used { + flex-basis: 100%; + font-size: 12px; + opacity: 0.82; +} + /* 充值套餐 */ .plans-section { padding: 0 20px 20px; -- 2.54.0 From 3f7ff9329ad45aca61bd33b5fdbf1f651b6ab890 Mon Sep 17 00:00:00 2001 From: stefanfeng Date: Tue, 25 Aug 2026 15:36:49 +0800 Subject: [PATCH 05/10] fix(avatar): align QA cards to the left --- digital-avatar-app/src/views/KnowledgeManage.vue | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/digital-avatar-app/src/views/KnowledgeManage.vue b/digital-avatar-app/src/views/KnowledgeManage.vue index f184885..0a38319 100644 --- a/digital-avatar-app/src/views/KnowledgeManage.vue +++ b/digital-avatar-app/src/views/KnowledgeManage.vue @@ -286,7 +286,12 @@ onMounted(async () => { .card-meta, .card-detail { margin: 5px 0 0; color: #9398AE; font-size: 11px; line-height: 1.4; }.card-detail { color: #8B6B58; } .card-delete { flex: 0 0 auto; align-self: center; border: 0; color: #EF4444; background: #FEF2F2; border-radius: 8px; padding: 7px 9px; font-size: 12px; cursor: pointer; } .card-empty { padding: 42px 16px; border: 1px dashed #F1D9C3; border-radius: 16px; color: #9398AE; background: #fff; font-size: 14px; text-align: center; } -.qa-card { align-items: stretch; }.qa-card.qa-disabled { opacity: .58; } +.qa-card { align-items: stretch; text-align: left; }.qa-card.qa-disabled { opacity: .58; } +.qa-card .card-content, +.qa-card .qa-question, +.qa-card .qa-answer, +.qa-card .card-meta, +.qa-card .qa-card-actions { text-align: left; } .qa-card-head { display: flex; align-items: center; justify-content: space-between; margin-bottom: 8px; }.qa-label { color: #C15F18; font-size: 11px; font-weight: 700; } .qa-question { display: block; color: #27201C; font-size: 15px; line-height: 1.5; }.qa-answer { display: -webkit-box; margin: 7px 0 0; overflow: hidden; color: #6B7280; font-size: 13px; line-height: 1.55; -webkit-box-orient: vertical; -webkit-line-clamp: 3; } .qa-card-actions { display: flex; gap: 8px; margin-top: 11px; } @@ -507,6 +512,8 @@ onMounted(async () => { .knowledge-card { display: grid; grid-template-columns: 42px minmax(0, 1fr); align-items: start; gap: 10px; padding: 13px; } .card-content { grid-column: 2; } .card-delete { grid-column: 2; justify-self: end; margin-top: -2px; } + .qa-card { display: block; } + .qa-card .card-content { width: 100%; grid-column: 1; } .card-title-row { align-items: flex-start; flex-wrap: wrap; gap: 5px 7px; } .status-pill { order: 2; } .search-bar { gap: 8px; } -- 2.54.0 From dc34a0335765d7243650b3647e34011f1a1d4ed8 Mon Sep 17 00:00:00 2001 From: stefanfeng Date: Tue, 25 Aug 2026 16:50:33 +0800 Subject: [PATCH 06/10] feat(ai): add dedicated digital avatar model config --- backend/app/api/endpoints/ai_models.py | 56 ++++++++++- backend/app/core/config.py | 1 + backend/app/core/database.py | 13 +++ backend/app/models/__init__.py | 1 + backend/app/schemas/__init__.py | 4 + backend/app/services/ai_service.py | 4 +- digital-avatar-app/backend/routers/chat.py | 61 +++++++----- .../backend/services/chat_model_config.py | 93 +++++++++++++++++++ .../backend/tests/test_chat_model_config.py | 79 ++++++++++++++++ digital-avatar-app/docker-compose.yml | 3 + docker-compose.yml | 1 + docker/mysql/init.sql | 1 + frontend/src/views/AIModels.vue | 20 +++- 13 files changed, 305 insertions(+), 32 deletions(-) create mode 100644 digital-avatar-app/backend/services/chat_model_config.py create mode 100644 digital-avatar-app/backend/tests/test_chat_model_config.py 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 }}
+ + {{ scopeLabels[m.usage_scope] || '通用业务' }} + 默认 禁用
@@ -49,6 +52,13 @@ + + + 通用业务 + 数字分身专用 + +
数字分身专用模型仅用于分身对话和主动接管回复
+
@@ -131,8 +141,9 @@ const testResult = ref(null) const testing = ref(false) const providerLabels = { openai: 'OpenAI', zhipu: '智谱GLM', wenxin: '文心一言', qianwen: '通义千问', local: '本地模型' } -const form = reactive({ model_name: '', provider: 'openai', api_base_url: '', api_key: '', model_version: '', temperature: 0.7, max_tokens: 1000, timeout_seconds: 30, is_default: 0 }) -const rules = { model_name: [{ required: true, message: '请输入模型名称' }], provider: [{ required: true }] } +const scopeLabels = { general: '通用业务', digital_avatar: '数字分身专用' } +const form = reactive({ model_name: '', provider: 'openai', usage_scope: 'general', api_base_url: '', api_key: '', model_version: '', temperature: 0.7, max_tokens: 1000, timeout_seconds: 30, is_default: 0 }) +const rules = { model_name: [{ required: true, message: '请输入模型名称' }], provider: [{ required: true }], usage_scope: [{ required: true }] } async function load() { const res = await getAIModels() @@ -155,13 +166,13 @@ function onProviderChange(provider) { function openCreate() { editModel.value = null - Object.assign(form, { model_name: '', provider: 'openai', api_base_url: PROVIDER_DEFAULTS.openai.api_base_url, api_key: '', model_version: PROVIDER_DEFAULTS.openai.model_version, temperature: 0.7, max_tokens: 1000, timeout_seconds: 30, is_default: 0 }) + Object.assign(form, { model_name: '', provider: 'openai', usage_scope: 'general', api_base_url: PROVIDER_DEFAULTS.openai.api_base_url, api_key: '', model_version: PROVIDER_DEFAULTS.openai.model_version, temperature: 0.7, max_tokens: 1000, timeout_seconds: 30, is_default: 0 }) dialogVisible.value = true } function openEdit(m) { editModel.value = m - Object.assign(form, { model_name: m.model_name, provider: m.provider, api_base_url: m.api_base_url || '', api_key: '', model_version: m.model_version || '', temperature: m.temperature, max_tokens: m.max_tokens, timeout_seconds: m.timeout_seconds, is_default: m.is_default }) + Object.assign(form, { model_name: m.model_name, provider: m.provider, usage_scope: m.usage_scope || 'general', api_base_url: m.api_base_url || '', api_key: '', model_version: m.model_version || '', temperature: m.temperature, max_tokens: m.max_tokens, timeout_seconds: m.timeout_seconds, is_default: m.is_default }) dialogVisible.value = true } @@ -236,4 +247,5 @@ onMounted(load) .result-meta { display: flex; align-items: center; gap: 10px; margin-bottom: 10px; } .result-content { background: var(--color-bg); border: 1px solid var(--color-border); border-radius: 8px; padding: 12px; font-size: 13px; line-height: 1.6; white-space: pre-wrap; max-height: 200px; overflow-y: auto; } .empty-state { grid-column: 1/-1; padding: 40px; } +.scope-tip { margin-top: 6px; color: var(--color-text-muted); font-size: 12px; line-height: 1.5; } -- 2.54.0 From e24e89d3269e263c7f290aad35e48e19da10252c Mon Sep 17 00:00:00 2001 From: stefanfeng Date: Tue, 25 Aug 2026 16:59:15 +0800 Subject: [PATCH 07/10] fix(deploy): serialize model migration and pin nginx --- backend/app/core/database.py | 24 ++++++++++++++---------- frontend/Dockerfile | 3 ++- 2 files changed, 16 insertions(+), 11 deletions(-) diff --git a/backend/app/core/database.py b/backend/app/core/database.py index ae3ccd6..80343ca 100755 --- a/backend/app/core/database.py +++ b/backend/app/core/database.py @@ -66,17 +66,21 @@ async def init_db(): 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" + await conn.execute(text("SELECT GET_LOCK('ai_model_usage_scope_migration', 30)")) + try: + 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'" )) - logger.info("AI模型配置表已增加 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 字段") + finally: + await conn.execute(text("SELECT RELEASE_LOCK('ai_model_usage_scope_migration')")) logger.info("✅ 数据库模型注册成功") logger.info("✅ 数据库初始化完成") diff --git a/frontend/Dockerfile b/frontend/Dockerfile index f3cd1dd..9b8af48 100644 --- a/frontend/Dockerfile +++ b/frontend/Dockerfile @@ -6,7 +6,8 @@ RUN npm install COPY . . RUN npm run build -FROM nginx:alpine +# Nginx 1.31 uses syscalls that are blocked by the test server's legacy kernel. +FROM nginx:1.28.3-alpine COPY --from=build /app/dist /usr/share/nginx/html COPY nginx.conf /etc/nginx/conf.d/default.conf EXPOSE 80 -- 2.54.0 From 4029c31ed736532bc363733d6794e8e20af4c124 Mon Sep 17 00:00:00 2001 From: stefanfeng Date: Tue, 25 Aug 2026 17:19:30 +0800 Subject: [PATCH 08/10] fix(avatar): bundle uni bridge and add favicon --- digital-avatar-app/index.html | 3 +-- digital-avatar-app/package-lock.json | 7 +++++++ digital-avatar-app/package.json | 1 + digital-avatar-app/public/favicon.svg | 11 +++++++++++ digital-avatar-app/src/main.ts | 4 ++++ digital-avatar-app/src/types/uni-webview-js.d.ts | 4 ++++ 6 files changed, 28 insertions(+), 2 deletions(-) create mode 100644 digital-avatar-app/public/favicon.svg create mode 100644 digital-avatar-app/src/types/uni-webview-js.d.ts diff --git a/digital-avatar-app/index.html b/digital-avatar-app/index.html index a9ce14e..1832bf8 100644 --- a/digital-avatar-app/index.html +++ b/digital-avatar-app/index.html @@ -7,8 +7,7 @@ content="width=device-width, initial-scale=1.0, maximum-scale=1.0, user-scalable=no, viewport-fit=cover" /> 会会数字分身 - - + -// 引入后全局会出现 window.uni.webView,H5 即可用 postMessage 与原生通信。 +// uni-webview bridge is bundled by main.ts; no external CDN is required. const BRIDGE_HANDLER = '__uniBridgeHandle__' @@ -15,6 +13,16 @@ export interface UniLaunchParams { ts?: string } +const PARAM_KEYS: (keyof UniLaunchParams)[] = ['token', 'userId', 'nickname', 'avatar', 'ts'] + +function readParams(search: string, target: UniLaunchParams): void { + const sp = new URLSearchParams(search) + for (const key of PARAM_KEYS) { + const value = sp.get(key) + if (value) target[key] = value + } +} + // 是否运行在 uniapp web-view 环境中 export function isInUniWebView(): boolean { return !!(window as any).uni?.webView @@ -22,21 +30,33 @@ export function isInUniWebView(): boolean { // 解析 web-view 加载 URL 时原生注入的参数(token / 会会用户) export function getLaunchParams(): UniLaunchParams { - const sp = new URLSearchParams(window.location.search) const params: UniLaunchParams = {} - const token = sp.get('token') - const userId = sp.get('userId') - const nickname = sp.get('nickname') - const avatar = sp.get('avatar') - const ts = sp.get('ts') - if (token) params.token = token - if (userId) params.userId = userId - if (nickname) params.nickname = decodeURIComponent(nickname) - if (avatar) params.avatar = decodeURIComponent(avatar) - if (ts) params.ts = ts + readParams(window.location.search, params) + const hashQueryIndex = window.location.hash.indexOf('?') + if (hashQueryIndex >= 0) { + readParams(window.location.hash.slice(hashQueryIndex + 1), params) + } return params } +// Remove the one-time login credential before any route is rendered or logged. +export function stripLaunchToken(): void { + const url = new URL(window.location.href) + url.searchParams.delete('token') + + const hash = url.hash.slice(1) + const queryIndex = hash.indexOf('?') + if (queryIndex >= 0) { + const path = hash.slice(0, queryIndex) + const hashParams = new URLSearchParams(hash.slice(queryIndex + 1)) + hashParams.delete('token') + const query = hashParams.toString() + url.hash = `${path}${query ? `?${query}` : ''}` + } + + window.history.replaceState(window.history.state, '', `${url.pathname}${url.search}${url.hash}`) +} + // H5 → 原生:发送事件(需引入 uniapp web-view bridge) export function postToNative(message: Record): boolean { if (!isInUniWebView()) return false diff --git a/digital-avatar-app/src/views/SmsLogin.vue b/digital-avatar-app/src/views/SmsLogin.vue index d5f7326..f49a0d2 100644 --- a/digital-avatar-app/src/views/SmsLogin.vue +++ b/digital-avatar-app/src/views/SmsLogin.vue @@ -138,7 +138,7 @@