diff --git a/AGENTS.md b/AGENTS.md index 73e1e53..32bdbfd 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -24,7 +24,7 @@ | 上传块 | `upload_block` | `{{{}}}` / `{{{upload file type:...}}}` | `plugins/uploadBlockPlugin.ts` | `utils/uploadBlock.js` | - **已验证**:当前代码完全符合"禁止嵌套、自动解析复原"的要求。修改功能块相关逻辑时需验证:1) schema 的 atom/isolating 属性不被移除;2) Remark/parseMarkdown/toMarkdown 解析链路完整。 -- 当前代码中可以确认的主功能是:AI 补全、OCR、文档转 Markdown、TTS、Markdown/DOCX/PDF 导入导出。 +- 当前代码中可以确认的主功能是:AI 补全、OCR、文档转 Markdown、TTS/ASR、Markdown/DOCX/PDF 导入导出。 - 历史文档中有一部分 TTS/ASR、Apple Silicon、Whisper、离线模式说明已经落后于当前代码;出现冲突时以实际代码和测试为准。 ## 先看哪里 @@ -80,8 +80,13 @@ - 文档超过 32 KB 时,AI 补全会在前端和插件层被禁用。 - OCR 文本和文档块内容会被注入补全上下文,但这些内容属于隐藏上下文,不应被直接当作用户可见文本重复输出。 - /v1/convert 当前支持 txt、docx、pptx、pdf,非 txt 文件通过 MarkItDown 转成 Markdown,之后会清理图片标记。 -- 前端存在 /v1/export/pdf 调用点,但当前后端主路由中看不到同名端点;排查 PDF 导出问题前先确认服务端是否真正提供该接口。 -- 当前 tts_asr.py 主要提供 TTS 相关能力。不要直接沿用 README 或历史修复文档里关于 ASR、Whisper、MPS/offline 的描述。 +- **AI 开关是全局广播状态**:MilkdownEditor.vue 通过 `llm-in-text:copilot-toggle` 同步主编辑器、文档块嵌套编辑器、网页搜索块嵌套编辑器;修 ghost text 时要同时检查这三处。 +- **设置项已从 currency 改为 country**:前后端请求、prompt、store、设置面板统一使用 `country`;仅在读取旧 localStorage 时兼容 `currency` 作为迁移兜底。 +- **DOCX/PDF 导出改为纯前端**:不再依赖 `/v1/export/pdf`。当前策略是先展开所有功能块,再从编辑器 HTML 构建导出内容;`src/utils/richExport.js` 负责 HTML -> PDF / DOCX。 +- **上传单文件限制统一为 100MB**:前端校验和后端 OCR 风控上限都按 100MB 处理。 +- **视频解析策略**:上传视频时,后端 `/v1/ocr` 接收 `media_type=video`,视频画面走 OCR 模型,音轨通过 ffmpeg 抽取后走 ASR 模型,最终合并为“视频画面 OCR + 视频音频 ASR”文本。 +- **OCR 明确关闭思考**:backend/llm.py 的 OCR payload 显式下发 `options.think = False` 与 `temperature = 0`。 +- **TTS/ASR 当前真实实现**:backend/tts_asr.py 已切到 `Qwen3TTSModel + faster-whisper` 路线;不要继续按旧的 MLX-only 文档理解当前实现。 ## 常用命令 @@ -119,7 +124,7 @@ 6. 再次用 `docker compose exec -T ...` 验证容器内文件和行为,不能只看本地文件。 - Docker 持久化数据统一落在部署目录内的 `docker-data/`,包括 PostgreSQL、Redis 和任务共享临时目录。 - 容器内访问宿主机模型服务时,不要继续使用 `localhost`;应改成 `host.docker.internal` 之类的容器可达地址。 -- 当前 Docker 部署默认使用轻量后端依赖集(`backend/requirements.docker.txt`),覆盖补全、OCR、转换、文档空间和队列,不默认包含本地 `torch` / TTS / ASR 模型栈。 +- 当前 Docker 部署的 `backend/requirements.docker.txt` 已包含 OCR、转换、队列以及 `torch` / `qwen-tts` / `faster-whisper`,并在 `backend/Dockerfile` 中额外安装 `ffmpeg` 以支持视频拆音轨。 - **Worker 容器**:worker.py 作为独立服务运行,通过 Redis Streams 消费任务队列。修改 job_handlers.py 或 worker.py 后需要验证 worker 容器内的代码已更新,可通过 `docker compose exec -T worker sh -lc "python -c 'from backend.job_handlers import get_handler; print(get_handler(\"completion\").__name__)'"` 验证。 - **Redis Streams 架构**:任务队列使用 Redis Streams,支持并发控制、速率限制和熔断器。job_system.py 定义 JOB_TYPES 和队列配置,worker.py 注册处理器并运行事件循环。 - 修改 Docker 相关文件时,除了代码本身,还要同步检查: diff --git a/backend/.env.example b/backend/.env.example index e89abbd..f52c5b1 100644 --- a/backend/.env.example +++ b/backend/.env.example @@ -103,7 +103,7 @@ RISK_WEB_SEARCH_MAX_OUTPUT_TOKENS=4096 RISK_WEB_SEARCH_TEMPERATURE=0.4 RISK_COMPRESS_MAX_INPUT_CHARS=128000 RISK_COMPRESS_MAX_OUTPUT_TOKENS=1536 -RISK_OCR_MAX_INPUT_BYTES=10485760 +RISK_OCR_MAX_INPUT_BYTES=104857600 # Web search providers SEARXNG_BASE_URL=http://searxng:8080 diff --git a/backend/AGENTS.md b/backend/AGENTS.md index 7fbc770..fd93f4a 100644 --- a/backend/AGENTS.md +++ b/backend/AGENTS.md @@ -4,7 +4,7 @@ ## 后端职责 -- 对外提供补全、取消补全、OCR、文档转换和 TTS 相关接口。 +- 对外提供补全、取消补全、OCR、文档转换和 TTS/ASR 相关接口。 - 组织 Prompt,上下文清洗,调用 Ollama 模型。 - **通过 Redis Streams 异步任务队列处理各类作业(completion/PRO/web_search/compress/OCR/convert/TTS/ASR)。** - 负责 API Key 校验、日志记录和部分启动预热逻辑。 @@ -62,10 +62,12 @@ ### /v1/ocr -- 把 base64 图片解码成字节。 -- 调用 call_vlm_ocr。 +- 把 base64 媒体内容解码成字节。 +- 支持 `media_type=image|video` 和 `mime_type`。 +- 图片直接调用 call_vlm_ocr。 +- 视频先把整段视频送入 OCR 模型,再用 ffmpeg 抽取音轨交给 ASR,最后合并文本。 - **结果通过 job_handlers.py ocr_handler 处理。** -- 返回识别文本和原始文件名。 +- 返回识别文本、原始文件名,以及视频场景下的 `ocr_text` / `asr_text`。 ### /v1/convert @@ -105,9 +107,9 @@ ### /v1/tts-asr/* - 通过 _register_tts_asr_routes 延迟导入并挂到主应用。 -- **当前代码里的 tts_asr.py 主要是 TTS 能力,不要自行假设存在完整 ASR 实现。** - **TTS 请求通过 job_handlers.py tts_handler 处理。** -- **支持多种语音模型(edge-tts、macos-say、pyttsx3)。** +- **ASR 请求通过 job_handlers.py asr_handler 处理。** +- **当前实现是 `Qwen3TTSModel + faster-whisper`,不是旧的 edge-tts / macos-say / MLX-only 路线。** ## 开发命令 @@ -146,9 +148,10 @@ - **验证码路由**:captcha_api.py 提供 /captcha/generate 和 /captcha/verify 端点,用于前端验证用户输入。 - **文档存储**:docs_store.py 提供 MIME 类型检测、文本/二进制分类和预览提取(8MB 限制)。 - **LLM 策略解析**:llm_policy.py 按 job_type 解析模型配置(模型名、温度、最大 token、思考级别)。 +- **OCR 明确关闭思考**:llm.py 的 `call_vlm_ocr` 对 OCR 请求显式设置 `options.think = False` 和 `temperature = 0`。 - **补全接口当前不是流式响应,不要按 SSE 方式改造周边代码。** - **ACTIVE_COMPLETIONS 在补全和取消路径里都被读写,任务生命周期要谨慎处理。** -- **main.py 里虽然有 _convert_docx_to_pdf 辅助函数,但当前 /v1/convert 路径实际走的是 MarkItDown,不要误以为 DOCX 转 PDF 桥接脚本已接入主流程。** +- **`/v1/export/pdf` 不是当前主链路**;DOCX/PDF 导出已经转到前端 `src/utils/richExport.js`。 - **API_KEY 存在占位默认值,这更像本地开发兜底,不是推荐的安全模式。** - **历史 TTS/ASR 文档和部分测试覆盖的是旧实现;代码与文档冲突时,先确认产品方向,再决定修代码还是修文档。** diff --git a/backend/Dockerfile b/backend/Dockerfile index 7be90a7..a144e47 100644 --- a/backend/Dockerfile +++ b/backend/Dockerfile @@ -6,6 +6,10 @@ ENV PYTHONUNBUFFERED=1 WORKDIR /app/backend +RUN apt-get update \ + && apt-get install -y --no-install-recommends ffmpeg \ + && rm -rf /var/lib/apt/lists/* + COPY backend/requirements.docker.txt /tmp/requirements.docker.txt RUN pip install --no-cache-dir -r /tmp/requirements.docker.txt diff --git a/backend/job_handlers.py b/backend/job_handlers.py index b7a09bd..564bc9e 100644 --- a/backend/job_handlers.py +++ b/backend/job_handlers.py @@ -14,6 +14,7 @@ import markitdown from audit_store import get_audit_store from llm import call_ollama, call_vlm_ocr, stream_ollama_events +from media_utils import extract_audio_wav_bytes, is_video_filename from prompt import ( build_completion_prompts, build_pro_completion_prompts, @@ -720,14 +721,60 @@ async def ocr_handler( path = payload["input_path"] identity, risk, lock_keys = await _enter_llm_execution(payload, emit) try: + filename = payload.get("filename", "image.jpg") + language = payload.get("language", "auto") + media_type = payload.get("media_type", "image") + mime_type = payload.get("mime_type", "") or "" + with open(path, "rb") as handle: - image_bytes = handle.read() - text = await call_vlm_ocr(image_bytes, payload.get("language", "auto")) + media_bytes = handle.read() + + await emit("progress", {"phase": "ocr", "media_type": media_type}) + ocr_text = await call_vlm_ocr( + media_bytes, + language, + mime_type=mime_type or "application/octet-stream", + media_type=media_type, + ) if is_cancelled(): raise asyncio.CancelledError() - await emit("result", {"text": text}) - await _exit_llm_execution(payload, identity, risk, lock_keys, status="completed", actual_output_text=text) - return {"text": text, "filename": payload.get("filename", "image.jpg")} + + result = { + "text": ocr_text, + "ocr_text": ocr_text, + "filename": filename, + "media_type": media_type, + } + + if media_type == "video" or is_video_filename(filename, mime_type): + asr_text = "" + if generate_asr_response is not None: + try: + await emit("progress", {"phase": "asr", "media_type": media_type}) + audio_bytes = await asyncio.to_thread(extract_audio_wav_bytes, path) + asr_response = await generate_asr_response(audio_bytes, language) + asr_text = getattr(asr_response, "text", "") or "" + except Exception as exc: + asr_text = f"(音频解析失败: {exc})" + if ocr_text.strip() or asr_text.strip(): + text_parts = [] + if ocr_text.strip(): + text_parts.append(f"## 视频画面 OCR\n\n{ocr_text.strip()}") + if asr_text.strip(): + text_parts.append(f"## 视频音频 ASR\n\n{asr_text.strip()}") + result["text"] = "\n\n".join(text_parts) + result["asr_text"] = asr_text + + await emit("result", result) + await _exit_llm_execution( + payload, + identity, + risk, + lock_keys, + status="completed", + actual_output_text=result["text"], + ) + return result except asyncio.CancelledError: await _exit_llm_execution(payload, identity, risk, lock_keys, status="cancelled", error_code="cancelled") raise diff --git a/backend/job_system.py b/backend/job_system.py index ccbac7a..640e04a 100644 --- a/backend/job_system.py +++ b/backend/job_system.py @@ -77,24 +77,6 @@ class QueueFullError(JobSystemError): self.max_queue = max_queue -# 修复:添加完整的异常类实现 -class JobSystemError(RuntimeError): - def __init__(self, message: str = "任务系统错误", error_code: str = "job_system_error") -> None: - super().__init__(message) - self.message = message - self.error_code = error_code - - def __str__(self) -> str: - return f"{self.error_code}: {self.message}" - - -class QueueFullError(JobSystemError): - def __init__(self, job_type: str, max_queue: int) -> None: - super().__init__(f"{job_type} 队列已满 (当前: {max_queue}/{max_queue})", "queue_full") - self.job_type = job_type - self.max_queue = max_queue - - @dataclass(frozen=True) class QueueConfig: job_type: str @@ -613,7 +595,14 @@ class RedisWorker: group = self.manager._group_name(job_type) semaphore = self.queue_semaphores[job_type] while True: - streams = await self.manager.redis.xreadgroup(group, self.consumer_name, {queue_key: ">"}, count=1, block=1000) + try: + streams = await self.manager.redis.xreadgroup(group, self.consumer_name, {queue_key: ">"}, count=1, block=1000) + except asyncio.CancelledError: + raise + except Exception as exc: + logger.warning("redis worker consume loop retrying type=%s error=%s", job_type, exc) + await asyncio.sleep(1) + continue if not streams: continue for _, messages in streams: diff --git a/backend/llm.py b/backend/llm.py index 75a01d2..3f5f440 100644 --- a/backend/llm.py +++ b/backend/llm.py @@ -551,31 +551,43 @@ async def stream_ollama_events( ) -async def call_vlm_ocr(image_bytes: bytes, language: str = 'auto') -> str: - """OCR via VLM using OpenAI-compatible vision API (image_url content part).""" +async def call_vlm_ocr( + media_bytes: bytes, + language: str = 'auto', + *, + mime_type: str = 'image/png', + media_type: str = 'image', +) -> str: + """OCR via VLM using OpenAI-compatible multimodal API.""" start = time.perf_counter() start_dt = datetime.now() logger.info( - '[VLM][ocr] request model=%s base_url=%s image_bytes=%d language=%s', - VLM_MODEL, LLM_BASE_URL, len(image_bytes), language, + '[VLM][ocr] request model=%s base_url=%s media_type=%s media_bytes=%d language=%s mime=%s', + VLM_MODEL, LLM_BASE_URL, media_type, len(media_bytes), language, mime_type, ) - image_b64 = base64.b64encode(image_bytes).decode('ascii') + media_b64 = base64.b64encode(media_bytes).decode('ascii') + content_part_type = 'video_url' if media_type == 'video' else 'image_url' + url_key = 'video_url' if media_type == 'video' else 'image_url' payload = { 'model': VLM_MODEL, 'messages': [{ 'role': 'user', 'content': [ - {'type': 'text', 'text': get_vlm_ocr_prompt()}, + {'type': 'text', 'text': f"{get_vlm_ocr_prompt()}\n\nTarget language hint: {language or 'auto'}"}, { - 'type': 'image_url', - 'image_url': {'url': f'data:image/png;base64,{image_b64}'}, + 'type': content_part_type, + url_key: {'url': f'data:{mime_type or "application/octet-stream"};base64,{media_b64}'}, }, ], }], 'stream': False, + 'options': { + 'temperature': 0, + 'think': False, + }, } http_timeout = httpx.Timeout(connect=10.0, read=None, write=30.0, pool=30.0) diff --git a/backend/main.py b/backend/main.py index 024241c..9720633 100644 --- a/backend/main.py +++ b/backend/main.py @@ -106,6 +106,8 @@ class OCRRequest(BaseModel): image: str filename: str = "image.jpg" language: str = "auto" + media_type: str = "image" + mime_type: str | None = None class ConvertRequest(BaseModel): @@ -281,8 +283,10 @@ def _register_handlers() -> None: 确保 Redis 重连或实例重建后处理器不会丢失。""" global _handlers_registered manager = get_job_manager() - # 强制清空旧 handlers,避免重复注册累积 - manager.handlers.clear() + # 强制清空旧 handlers,避免重复注册累积;测试替身只暴露 register_handler。 + handlers = getattr(manager, "handlers", None) + if handlers is not None: + handlers.clear() manager.register_handler("completion", completion_handler) manager.register_handler("pro_completion", pro_completion_handler) manager.register_handler("web_search", web_search_handler) @@ -294,7 +298,8 @@ def _register_handlers() -> None: _handlers_registered = True # 打印注册信息便于调试 - logger.info("handlers registered: %s", list(manager.handlers.keys()))tered: %s", list(manager.handlers.keys())) + registered = list(getattr(manager, "handlers", {}).keys()) + logger.info("handlers registered: %s", registered) def _sse(event: str, data: dict) -> str: @@ -623,16 +628,28 @@ async def ocr_image(request: Request, req: OCRRequest, auth: dict = Security(_au except Exception as exc: return JSONResponse({"error": str(exc)}, status_code=500) if len(image_bytes) > config.ocr_max_input_bytes: - return JSONResponse({"error": "图片过大,无法执行 OCR"}, status_code=400) + return JSONResponse({"error": "文件过大,无法执行 OCR/视频解析"}, status_code=400) input_path = persist_temp_input(image_bytes, os.path.splitext(req.filename)[1] or ".img") try: identity, payload = await _prepare_llm_payload( request, job_type="ocr", - request_body={"filename": req.filename, "language": req.language, "image_bytes": len(image_bytes)}, + request_body={ + "filename": req.filename, + "language": req.language, + "media_type": req.media_type, + "mime_type": req.mime_type, + "image_bytes": len(image_bytes), + }, raw_size=len(image_bytes), - token_source_text=f"{req.filename}:{len(image_bytes)}:{req.language}", - extra_payload={"input_path": input_path, "filename": req.filename, "language": req.language}, + token_source_text=f"{req.filename}:{req.media_type}:{req.mime_type}:{len(image_bytes)}:{req.language}", + extra_payload={ + "input_path": input_path, + "filename": req.filename, + "language": req.language, + "media_type": req.media_type, + "mime_type": req.mime_type, + }, ) job_id = await _queue_job("ocr", payload, identity.request_id) except RiskRejected as exc: diff --git a/backend/media_utils.py b/backend/media_utils.py new file mode 100644 index 0000000..e9dd176 --- /dev/null +++ b/backend/media_utils.py @@ -0,0 +1,47 @@ +from __future__ import annotations + +import os +import subprocess +import tempfile + + +VIDEO_EXTENSIONS = {".mp4", ".webm", ".mov", ".avi", ".mkv", ".m4v", ".ogv"} + + +def is_video_filename(filename: str = "", mime_type: str = "") -> bool: + ext = os.path.splitext(filename or "")[1].lower() + mime = (mime_type or "").strip().lower() + return ext in VIDEO_EXTENSIONS or mime.startswith("video/") + + +def extract_audio_wav_bytes(input_path: str) -> bytes: + if not input_path or not os.path.exists(input_path): + raise FileNotFoundError("输入媒体文件不存在") + + fd, output_path = tempfile.mkstemp(suffix=".wav") + os.close(fd) + try: + subprocess.run( + [ + "ffmpeg", + "-y", + "-i", + input_path, + "-vn", + "-acodec", + "pcm_s16le", + "-ar", + "16000", + "-ac", + "1", + output_path, + ], + check=True, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + ) + with open(output_path, "rb") as handle: + return handle.read() + finally: + if os.path.exists(output_path): + os.unlink(output_path) diff --git a/backend/models.py b/backend/models.py index adef2a1..1a737c2 100644 --- a/backend/models.py +++ b/backend/models.py @@ -5,5 +5,5 @@ from pydantic import BaseModel class UserPreferences(BaseModel): """用户偏好设置""" language: str = "auto" - currency: str = "auto" + country: str = "auto" timezone: str = "auto" diff --git a/backend/prompt.py b/backend/prompt.py index 4282e12..1a14e08 100644 --- a/backend/prompt.py +++ b/backend/prompt.py @@ -355,8 +355,8 @@ def build_completion_prompts( if preferences: if preferences.language and preferences.language != "auto": pref_info.append(f"Preferred language: {preferences.language}") - if preferences.currency and preferences.currency != "auto": - pref_info.append(f"Preferred currency: {preferences.currency}") + if preferences.country and preferences.country != "auto": + pref_info.append(f"Preferred country: {preferences.country}") preferences_instruction = "\n".join(pref_info) if preferences_instruction: @@ -454,8 +454,8 @@ def build_pro_completion_prompts( if preferences: if preferences.language and preferences.language != "auto": pref_info.append(f"Preferred language: {preferences.language}") - if preferences.currency and preferences.currency != "auto": - pref_info.append(f"Preferred currency: {preferences.currency}") + if preferences.country and preferences.country != "auto": + pref_info.append(f"Preferred country: {preferences.country}") if preferences.timezone and preferences.timezone != "auto": pref_info.append(f"Preferred timezone: {preferences.timezone}") diff --git a/backend/requirements.docker.txt b/backend/requirements.docker.txt index 4d98848..3bb1fd6 100644 --- a/backend/requirements.docker.txt +++ b/backend/requirements.docker.txt @@ -8,3 +8,10 @@ python-multipart>=0.0.9 python-dotenv>=1.0.0 markitdown>=0.1.1 geoip2>=4.8.0 +numpy>=1.26.0 +torch>=2.2.0 +soundfile>=0.12.1 +scipy>=1.13.0 +qwen-tts +modelscope>=1.18.0 +faster-whisper>=1.1.0 diff --git a/backend/risk_config.py b/backend/risk_config.py index c781305..e6d3940 100644 --- a/backend/risk_config.py +++ b/backend/risk_config.py @@ -92,7 +92,7 @@ class RiskConfig: def load_risk_config() -> RiskConfig: raw_origins = _str_env( "CORS_ALLOW_ORIGINS", - "https://chat.imageteach.tech,http://localhost:8080,http://127.0.0.1:8080", + "https://chat.imageteach.tech,http://localhost:8080,http://127.0.0.1:8080,http://localhost:5173,http://127.0.0.1:5173", ) cors_allow_origins = tuple( origin.strip() for origin in raw_origins.split(",") if origin.strip() @@ -138,7 +138,7 @@ def load_risk_config() -> RiskConfig: web_search_temperature=_float_env("RISK_WEB_SEARCH_TEMPERATURE", 0.4), compress_max_input_chars=_int_env("RISK_COMPRESS_MAX_INPUT_CHARS", 128000), compress_max_output_tokens=_int_env("RISK_COMPRESS_MAX_OUTPUT_TOKENS", 1536), - ocr_max_input_bytes=_int_env("RISK_OCR_MAX_INPUT_BYTES", 10 * 1024 * 1024), + ocr_max_input_bytes=_int_env("RISK_OCR_MAX_INPUT_BYTES", 100 * 1024 * 1024), completion_input_cost_per_1k=_float_env("RISK_COMPLETION_INPUT_COST_PER_1K", 0.0004), completion_output_cost_per_1k=_float_env("RISK_COMPLETION_OUTPUT_COST_PER_1K", 0.0016), pro_input_cost_per_1k=_float_env("RISK_PRO_INPUT_COST_PER_1K", 0.003), diff --git a/backend/tests/test_llm.py b/backend/tests/test_llm.py index bdc6cfa..ea82fe6 100644 --- a/backend/tests/test_llm.py +++ b/backend/tests/test_llm.py @@ -299,3 +299,4 @@ def test_call_vlm_ocr(monkeypatch): image_part = [p for p in content_parts if p.get("type") == "image_url"] assert len(image_part) == 1 assert image_part[0]["image_url"]["url"].startswith("data:image/png;base64,") + assert captured["json"]["options"]["think"] is False diff --git a/backend/tests/test_main_endpoints.py b/backend/tests/test_main_endpoints.py index 5a767ca..c707d21 100644 --- a/backend/tests/test_main_endpoints.py +++ b/backend/tests/test_main_endpoints.py @@ -4,6 +4,7 @@ import importlib import os import sys from pathlib import Path +from types import SimpleNamespace from fastapi.testclient import TestClient @@ -158,6 +159,35 @@ def test_post_ocr_mocked(monkeypatch): assert "OCR result text" in body +def test_post_video_ocr_merges_ocr_and_asr(monkeypatch): + async def fake_ocr(*args, **kwargs): + return "画面文字" + + async def fake_asr(*args, **kwargs): + return SimpleNamespace(text="音频转写") + + monkeypatch.setattr(job_handlers, "call_vlm_ocr", fake_ocr) + monkeypatch.setattr(job_handlers, "generate_asr_response", fake_asr) + monkeypatch.setattr(job_handlers, "extract_audio_wav_bytes", lambda _path: b"fake wav") + + video_b64 = base64.b64encode(b"pretend video data").decode() + with TestClient(main.app) as client: + with client.stream("POST", "/v1/ocr", headers=HEADERS, json={ + "image": video_b64, + "filename": "sample.mp4", + "language": "auto", + "media_type": "video", + "mime_type": "video/mp4", + }) as resp: + assert resp.status_code == 200 + body = "".join(resp.iter_text()) + + assert "视频画面 OCR" in body + assert "视频音频 ASR" in body + assert "画面文字" in body + assert "音频转写" in body + + def test_post_convert_txt_returns_markdown(): content = base64.b64encode(b"hello world").decode() with TestClient(main.app) as client: diff --git a/backend/tests/test_pro_completions.py b/backend/tests/test_pro_completions.py index 20d1806..65107c5 100644 --- a/backend/tests/test_pro_completions.py +++ b/backend/tests/test_pro_completions.py @@ -35,7 +35,7 @@ def _payload(): "privacy_mode": True, "user_preferences": { "language": "zh", - "currency": "CNY", + "country": "CN", "timezone": "Asia/Shanghai", }, } @@ -43,7 +43,7 @@ def _payload(): def test_pro_queue_full_returns_429(monkeypatch): async def fake_queue_job(*args, **kwargs): - raise job_system.QueueFullError("pro_completion queue is full") + raise job_system.QueueFullError("pro_completion", 8) monkeypatch.setattr(main, "_queue_job", fake_queue_job) with TestClient(main.app) as client: @@ -79,13 +79,13 @@ def test_pro_prompt_accepts_serialized_preferences(): instruction="expand", preferences={ "language": "zh", - "currency": "CNY", + "country": "CN", "timezone": "Asia/Shanghai", }, ) assert "Preferred language: zh" in user_prompt - assert "Preferred currency: CNY" in user_prompt + assert "Preferred country: CN" in user_prompt assert "Preferred timezone: Asia/Shanghai" in user_prompt diff --git a/backend/tests/test_web_search.py b/backend/tests/test_web_search.py index 478d632..bd5a413 100644 --- a/backend/tests/test_web_search.py +++ b/backend/tests/test_web_search.py @@ -17,6 +17,7 @@ if str(BACKEND_DIR) not in sys.path: import job_handlers # type: ignore import job_system # type: ignore +import llm # type: ignore import risk_control # type: ignore import session_store # type: ignore import audit_store # type: ignore @@ -55,10 +56,14 @@ def test_web_search_route_returns_done(monkeypatch): return {"content": '["vector database comparison", "pinecone weaviate qdrant"]'} if tag.endswith("-webu"): return {"content": '["https://example.com/a", "https://example.com/b"]'} - if tag.endswith("-webf"): - return {"content": "第一段\n\n第二段"} raise AssertionError(f"unexpected tag: {tag}") + async def fake_stream_ollama_events(prompt, system_prompt=None, tag="", **kwargs): # noqa: ARG001 + if not tag.endswith("-webf"): + raise AssertionError(f"unexpected stream tag: {tag}") + yield "content", "第一段\n\n" + yield "content", "第二段" + async def fake_searxng_search(query, *, limit): # noqa: ARG001 return [ { @@ -83,6 +88,7 @@ def test_web_search_route_returns_done(monkeypatch): monkeypatch.setattr(job_handlers, "call_ollama", fake_call_ollama) monkeypatch.setattr(job_handlers, "_searxng_search", fake_searxng_search) monkeypatch.setattr(job_handlers, "_firecrawl_scrape", fake_firecrawl_scrape) + monkeypatch.setattr(llm, "stream_ollama_events", fake_stream_ollama_events) with TestClient(main.app) as client: with client.stream("POST", "/v1/web-search", headers=HEADERS, json=_payload()) as resp: diff --git a/backend/tts_asr.py b/backend/tts_asr.py index f77a67e..359332a 100644 --- a/backend/tts_asr.py +++ b/backend/tts_asr.py @@ -1,260 +1,133 @@ import asyncio import base64 -import io import logging import os import tempfile -import wave from typing import Optional -# 设置 Hugging Face / ModelScope 镜像源为国内镜像 os.environ.setdefault("HF_ENDPOINT", "https://hf-mirror.com") -import numpy as np -import torch from fastapi import APIRouter, HTTPException from pydantic import BaseModel logger = logging.getLogger(__name__) try: - asyncio.get_running_loop() -except RuntimeError: - asyncio.set_event_loop(asyncio.new_event_loop()) + import numpy as np # type: ignore +except Exception as exc: # pragma: no cover + logger.debug("numpy import failed: %s", exc) + np = None # type: ignore + +try: + import torch # type: ignore +except Exception as exc: # pragma: no cover + logger.debug("torch import failed: %s", exc) + torch = None # type: ignore -# New TTS model import try: from qwen_tts import Qwen3TTSModel # type: ignore -except Exception as e: # pragma: no cover - logger.debug("qwen_tts import failed (optional): %s", e) +except Exception as exc: # pragma: no cover + logger.debug("qwen_tts import failed: %s", exc) Qwen3TTSModel = None # type: ignore -# ASR model import (MLX-based, Apple Silicon only) try: - from mlx_audio.stt.models.qwen3_asr import ( # type: ignore - ForcedAlignerModel, - Qwen3ASRModel, - ) -except Exception as e: # pragma: no cover - logger.debug("mlx_audio import failed (optional): %s", e) - Qwen3ASRModel = None # type: ignore - ForcedAlignerModel = None # type: ignore + from faster_whisper import WhisperModel # type: ignore +except Exception as exc: # pragma: no cover + logger.debug("faster_whisper import failed: %s", exc) + WhisperModel = None # type: ignore try: from modelscope import snapshot_download # type: ignore -except Exception as e: # pragma: no cover - logger.debug("modelscope import failed (optional): %s", e) +except Exception as exc: # pragma: no cover + logger.debug("modelscope import failed: %s", exc) + snapshot_download = None # type: ignore meta_router = APIRouter() generation_router = APIRouter() -# Global model instances -_tts_model: Optional["Qwen3TTSModel"] = None -_asr_model: Optional[object] = None # Qwen3ASRModel or ForcedAlignerModel -_align_model: Optional[object] = None # Qwen3-ForcedAlignerModel - -# Model paths for loading MODEL_ID_HF = "Qwen/Qwen3-TTS-12Hz-1.7B-VoiceDesign" MODEL_ID_MS = "Qwen/Qwen3-TTS-12Hz-1.7B-VoiceDesign" +ASR_MODEL_ID = os.getenv("ASR_MODEL_ID", "small") +ASR_COMPUTE_TYPE = os.getenv("ASR_COMPUTE_TYPE", "int8") -# ModelScope ASR/ForcedAligner models (MLX 4-bit format) -ASR_MODEL_ID_MS = "aufklarer/Qwen3-ASR-0.6B-MLX-4bit" -ALIGN_MODEL_ID_MS = "aufklarer/Qwen3-ForcedAligner-0.6B-MLX" +_tts_model: Optional["Qwen3TTSModel"] = None +_asr_model: Optional["WhisperModel"] = None def _get_device_map() -> str: - """设备检测逻辑:优先 CUDA,其次 MPS,最后 CPU""" + if torch is None: + return "cpu" if torch.cuda.is_available(): - return "cuda:0" + return "cuda" try: if hasattr(torch.backends, "mps") and torch.backends.mps.is_available(): return "mps" - except Exception as e: # noqa: ANN001 - logger.debug("MPS check failed: %s", e) + except Exception as exc: # pragma: no cover + logger.debug("MPS check failed: %s", exc) return "cpu" -def _download_model_from_modelscope() -> Optional[str]: - """从 ModelScope 下载模型到本地缓存目录""" +def _download_tts_model_from_modelscope() -> Optional[str]: + if snapshot_download is None: + return None + cache_dir = os.path.join(os.path.dirname(__file__), "models") + os.makedirs(cache_dir, exist_ok=True) try: - cache_dir = os.path.join(os.path.dirname(__file__), "models") - os.makedirs(cache_dir, exist_ok=True) - model_dir = snapshot_download( - MODEL_ID_MS, - cache_dir=cache_dir, - revision="master" - ) - logger.info("ModelScope 模型下载完成: %s", model_dir) - return model_dir - except Exception as e: # noqa: ANN001 - logger.warning("ModelScope 下载失败: %s", e) + return snapshot_download(MODEL_ID_MS, cache_dir=cache_dir, revision="master") + except Exception as exc: # pragma: no cover + logger.warning("ModelScope TTS download failed: %s", exc) return None -async def _warmup_tts(): - """预热 TTS 模型""" - await asyncio.to_thread(_load_tts_model_with_retry) - - -async def _warmup_asr(): - """预热 ASR 模型(从 ModelScope 下载并加载)""" - await asyncio.to_thread(_load_asr_models) - - -async def _warmup_all(): - """预热所有模型(TTS 和 ASR)""" - logger.info("[Warmup] 开始预热 TTS 模型...") - await _warmup_tts() - logger.info("[Warmup] TTS 模型预热完成") - - if Qwen3ASRModel is not None: - logger.info("[Warmup] 开始预热 ASR 模型...") - await _warmup_asr() - logger.info("[Warmup] ASR 模型预热完成") - - -def _load_tts_model_with_retry(max_retries: int = 3) -> "Qwen3TTSModel": - """加载 TTS 模型,支持多个镜像源""" +def _ensure_tts_model() -> "Qwen3TTSModel": global _tts_model if _tts_model is not None: return _tts_model - if Qwen3TTSModel is None: - raise RuntimeError("qwen_tts 库未安装,无法加载 TTS 模型") + if np is None or torch is None or Qwen3TTSModel is None: + raise RuntimeError("TTS 依赖未安装完整") device_map = _get_device_map() - last_err = None + dtype = torch.float16 if device_map != "cpu" else torch.float32 - # 策略1: 尝试从 ModelScope 下载后加载 - for attempt in range(max_retries): - try: - logger.info("尝试从 ModelScope 下载 TTS 模型...") - model_path = _download_model_from_modelscope() - if model_path and os.path.isdir(model_path): - _tts_model = Qwen3TTSModel.from_pretrained( # type: ignore - model_path, - device_map=device_map, - dtype=torch.float16, - ) - logger.info("ModelScope TTS 模型加载成功: %s", model_path) - return _tts_model - except Exception as e: # noqa: ANN001 - logger.warning("ModelScope TTS 加载失败 (尝试 %d/%d): %s", attempt + 1, max_retries, e) - last_err = e + model_path = _download_tts_model_from_modelscope() + last_error = None - # 策略2: 尝试从 HuggingFace 镜像加载 - for attempt in range(max_retries): + for candidate in [model_path, MODEL_ID_HF]: + if not candidate: + continue try: - logger.info("尝试从 HuggingFace 镜像加载 TTS...") _tts_model = Qwen3TTSModel.from_pretrained( # type: ignore - MODEL_ID_HF, + candidate, device_map=device_map, - dtype=torch.float16, + dtype=dtype, ) - logger.info("HuggingFace TTS 模型加载成功") return _tts_model - except Exception as e: # noqa: ANN001 - logger.warning("HuggingFace TTS 加载失败 (尝试 %d/%d): %s", attempt + 1, max_retries, e) - last_err = e + except Exception as exc: + last_error = exc + logger.warning("TTS model load failed from %s: %s", candidate, exc) - raise RuntimeError(f"无法加载 TTS 模型: {last_err}") from last_err + raise RuntimeError(f"TTS 模型加载失败: {last_error}") from last_error -def _load_asr_models() -> None: - """从 ModelScope 下载并加载 ASR/ForcedAligner MLX 模型""" - global _asr_model, _align_model - - if snapshot_download is None: - logger.warning("modelscope 未安装,跳过 ASR 模型加载") - return - - if Qwen3ASRModel is None: - logger.warning("mlx_audio 未安装,跳过 ASR 模型加载") - return - - # Download and load ASR model from ModelScope - try: - logger.info("从 ModelScope 下载 ASR 模型...") - asr_cache_dir = os.path.join(os.path.dirname(__file__), "models", "asr") - asr_model_dir = snapshot_download(ASR_MODEL_ID_MS, cache_dir=asr_cache_dir) - _load_asr_from_path(asr_model_dir) - except Exception as e: # noqa: ANN001 - logger.warning("ASR ModelScope 下载失败,尝试 hf-mirror: %s", e) - try: - _load_asr_from_hf_mirror() - except Exception as e2: # noqa: ANN001 - logger.warning("ASR hf-mirror 加载失败,跳过 ASR: %s", e2) - - # Download and load ForcedAligner model from ModelScope - try: - logger.info("从 ModelScope 下载 ForcedAligner 模型...") - align_cache_dir = os.path.join(os.path.dirname(__file__), "models", "aligner") - align_model_dir = snapshot_download(ALIGN_MODEL_ID_MS, cache_dir=align_cache_dir) - _load_align_from_path(align_model_dir) - except Exception as e: # noqa: ANN001 - logger.warning("ForcedAligner ModelScope 下载失败,尝试 hf-mirror: %s", e) - try: - _load_align_from_hf_mirror() - except Exception as e2: # noqa: ANN001 - logger.warning("ForcedAligner hf-mirror 加载失败,跳过: %s", e2) - - -def _load_asr_from_path(model_dir: str) -> None: - """从本地路径加载 ASR MLX 模型""" +def _ensure_asr_model() -> "WhisperModel": global _asr_model - try: - from mlx_audio.stt.utils import load as stt_load # type: ignore + if _asr_model is not None: + return _asr_model + if WhisperModel is None: + raise RuntimeError("faster-whisper 未安装") - model = stt_load(model_dir) - _asr_model = model - logger.info("ASR 模型加载成功 (路径: %s)", model_dir) - except Exception as e: # noqa: ANN001 - logger.warning("ASR MLX 加载失败,尝试直接构建: %s", e) - try: - from mlx.core import load as mx_load # type: ignore - - weights = mx_load(os.path.join(model_dir, "model.safetensors")) - from mlx_lm import load as lm_load # type: ignore - - model = lm_load(model_dir, model_cls=Qwen3ASRModel) - _asr_model = model - except Exception as e2: # noqa: ANN001 - raise RuntimeError(f"无法加载 ASR MLX 模型: {e2}") from e + device = "cuda" if _get_device_map() == "cuda" else "cpu" + compute_type = ASR_COMPUTE_TYPE if device == "cpu" else "float16" + _asr_model = WhisperModel(ASR_MODEL_ID, device=device, compute_type=compute_type) + return _asr_model -def _load_asr_from_hf_mirror() -> None: - """从 hf-mirror 加载 ASR MLX 模型""" - global _asr_model - try: - from mlx_audio.stt.utils import load as stt_load # type: ignore - - model = stt_load("mlx-community/Qwen3-ASR-0.6B-4bit") - _asr_model = model - except Exception as e: # noqa: ANN001 - raise RuntimeError(f"无法从 hf-mirror 加载 ASR MLX: {e}") from e +async def _warmup_tts(): + await asyncio.to_thread(_ensure_tts_model) -def _load_align_from_path(model_dir: str) -> None: - """从本地路径加载 ForcedAligner MLX 模型""" - global _align_model - try: - from mlx_audio.stt.utils import load as stt_load # type: ignore - - model = stt_load(model_dir) - _align_model = model - except Exception as e: # noqa: ANN001 - raise RuntimeError(f"无法加载 ForcedAligner MLX 模型 (路径: {model_dir}): {e}") from e - - -def _load_align_from_hf_mirror() -> None: - """从 hf-mirror 加载 ForcedAligner MLX 模型""" - global _align_model - try: - from mlx_audio.stt.utils import load as stt_load # type: ignore - - model = stt_load("mlx-community/Qwen3-ForcedAligner-0.6B-4bit") - _align_model = model - except Exception as e: # noqa: ANN001 - raise RuntimeError(f"无法从 hf-mirror 加载 ForcedAligner MLX: {e}") from e +async def _warmup_asr(): + await asyncio.to_thread(_ensure_asr_model) class TTSRequest(BaseModel): @@ -286,43 +159,25 @@ class ModelStatus(BaseModel): device: str -def _ensure_tts_model() -> "Qwen3TTSModel": - """确保 TTS 模型已加载""" - global _tts_model - if _tts_model is None: - _tts_model = _load_tts_model_with_retry() - return _tts_model - - -def _ensure_asr_model(): - """确保 ASR 模型已加载(懒加载)""" - global _asr_model - if _asr_model is None: - try: - from mlx_audio.stt.utils import load as stt_load # type: ignore - - _asr_model = stt_load(ASR_MODEL_ID_MS) - except Exception as e: # noqa: ANN001 - raise RuntimeError(f"无法加载 ASR MLX 模型 (路径: {ASR_MODEL_ID_MS}): {e}") from e - return _asr_model - - -def _ensure_align_model(): - """确保 ForcedAligner 模型已加载(懒加载)""" - global _align_model - if _align_model is None: - try: - from mlx_audio.stt.utils import load as stt_load # type: ignore - - _align_model = stt_load(ALIGN_MODEL_ID_MS) - except Exception as e: # noqa: ANN001 - raise RuntimeError(f"无法加载 ForcedAligner MLX 模型 (路径: {ALIGN_MODEL_ID_MS}): {e}") from e - return _align_model +def _normalize_language(language: Optional[str]) -> Optional[str]: + if not language: + return None + value = language.strip().lower() + if value in {"auto", ""}: + return None + mapping = { + "zh-cn": "zh", + "zh-hans": "zh", + "zh-tw": "zh", + "en-us": "en", + "ja-jp": "ja", + "ko-kr": "ko", + } + return mapping.get(value, value.split("-")[0]) @meta_router.get("/status", response_model=ModelStatus) async def get_status(): - """获取模型状态""" return ModelStatus( tts_loaded=_tts_model is not None, asr_loaded=_asr_model is not None, @@ -332,11 +187,10 @@ async def get_status(): @meta_router.get("/config") async def get_config(): - """获取配置信息""" return { "model": { "tts": MODEL_ID_MS, - "asr": ASR_MODEL_ID_MS if Qwen3ASRModel is not None else None, + "asr": ASR_MODEL_ID, }, "device": _get_device_map(), "status": { @@ -348,15 +202,11 @@ async def get_config(): @meta_router.post("/warmup") async def warmup_models(): - """手动触发模型预热""" await _warmup_tts() - - if Qwen3ASRModel is not None: - await _warmup_asr() - + await _warmup_asr() return { "tts_warmup": _tts_model is not None, - "asr_warmup": _asr_model is not None if Qwen3ASRModel else False, + "asr_warmup": _asr_model is not None, "device": _get_device_map(), } @@ -367,26 +217,30 @@ async def generate_tts_response( speaker: str = "Vivian", output_format: str = "wav", ) -> TTSResponse: - del speaker # current model path does not expose multi-speaker routing - del output_format # current implementation always returns wav - try: - model = _ensure_tts_model() - except Exception as e: # noqa: ANN001 - raise HTTPException(status_code=500, detail=str(e)) + del speaker + del output_format + if np is None: + raise HTTPException(status_code=501, detail="numpy 未安装,TTS 功能不可用") try: - wavs, sr = model.generate_voice_design( # type: ignore + model = _ensure_tts_model() + except Exception as exc: + raise HTTPException(status_code=500, detail=str(exc)) + + try: + wavs, sample_rate = await asyncio.to_thread( + model.generate_voice_design, # type: ignore text=text, language="Chinese", instruct=instruct or "", ) - except Exception as e: # noqa: ANN001 - logger.exception("TTS 推理失败") - raise HTTPException(status_code=500, detail=f"TTS 推理失败: {e}") + except Exception as exc: + logger.exception("TTS inference failed") + raise HTTPException(status_code=500, detail=f"TTS 推理失败: {exc}") wav_data = wavs[0] if isinstance(wavs, (list, tuple)) else wavs - if hasattr(wav_data, 'numpy'): # type: ignore - wav_data = wav_data.cpu().numpy() # type: ignore + if hasattr(wav_data, "cpu"): + wav_data = wav_data.cpu().numpy() wav_data = np.asarray(wav_data, dtype=np.float32) tmp_path = None @@ -394,81 +248,66 @@ async def generate_tts_response( import soundfile as sf # type: ignore fd, tmp_path = tempfile.mkstemp(suffix=".wav") - os.close(fd) # type: ignore - sf.write(tmp_path, wav_data, sr) - with open(tmp_path, "rb") as f: # noqa: SIM115 - audio_bytes = f.read() - except Exception as e: # noqa: ANN001 - logger.exception("音频编码失败") - raise HTTPException(status_code=500, detail=f"音频编码失败: {e}") + os.close(fd) + sf.write(tmp_path, wav_data, sample_rate) + with open(tmp_path, "rb") as handle: + audio_bytes = handle.read() + except Exception as exc: + logger.exception("TTS audio encode failed") + raise HTTPException(status_code=500, detail=f"音频编码失败: {exc}") finally: - if tmp_path and os.path.exists(tmp_path): # noqa: SIM201 - try: - os.unlink(tmp_path) - except Exception: - pass + if tmp_path and os.path.exists(tmp_path): + os.unlink(tmp_path) - duration_ms = int(len(wav_data) / sr * 1000) if sr > 0 else 0 - audio_base64 = base64.b64encode(audio_bytes).decode("utf-8") - return TTSResponse(audio_base64=audio_base64, format="wav", duration_ms=duration_ms) + duration_ms = int(len(wav_data) / sample_rate * 1000) if sample_rate > 0 else 0 + return TTSResponse( + audio_base64=base64.b64encode(audio_bytes).decode("utf-8"), + format="wav", + duration_ms=duration_ms, + ) async def generate_asr_response(audio_bytes: bytes, language: Optional[str] = "zh-CN") -> ASRResponse: - if Qwen3ASRModel is None: - raise HTTPException(status_code=501, detail="mlx_audio 未安装,ASR 功能不可用") + if not audio_bytes: + raise HTTPException(status_code=400, detail="音频内容为空") try: model = _ensure_asr_model() - except Exception as e: # noqa: ANN001 - raise HTTPException(status_code=500, detail=f"ASR 模型加载失败: {e}") + except Exception as exc: + raise HTTPException(status_code=500, detail=f"ASR 模型加载失败: {exc}") + normalized_language = _normalize_language(language) + tmp_path = None try: - wav_buffer = io.BytesIO(audio_bytes) - with wave.open(wav_buffer, 'rb') as wf: # noqa: SIM115 - n_channels = wf.getnchannels() - sampwidth = wf.getsampwidth() - framerate = wf.getframerate() - n_frames = wf.getnframes() + fd, tmp_path = tempfile.mkstemp(suffix=".wav") + os.close(fd) + with open(tmp_path, "wb") as handle: + handle.write(audio_bytes) - raw_data = wf.readframes(n_frames) - audio_array = np.frombuffer(raw_data, dtype=np.int16 if sampwidth == 2 else np.float32) - - if n_channels > 1: - audio_array = np.mean(audio_array.reshape(-1, n_channels), axis=1) - - if framerate != 16000: - try: - import scipy.signal as signal # type: ignore - - n_samples = int(len(audio_array) * 16000 / framerate) - audio_array = signal.resample(audio_array, n_samples) # type: ignore - except Exception as e2: # noqa: ANN001 - logger.warning("重采样失败,使用原始音频: %s", e2) - - if audio_array.dtype == np.int16: - audio_array = audio_array.astype(np.float32) / 32768.0 - - result = model.generate( # type: ignore - audio_array, - language=language if language else None, + segments, info = await asyncio.to_thread( + model.transcribe, + tmp_path, + language=normalized_language, + vad_filter=True, + beam_size=5, ) - - recognized_text = getattr(result, 'text', str(result)) if hasattr(result, 'text') else str(result) - detected_lang = getattr(result, 'language', language or "zh-CN") - if isinstance(detected_lang, list) and len(detected_lang) > 0: - detected_lang = detected_lang[0] - - return ASRResponse(text=recognized_text, language=str(detected_lang)) + text = "".join(segment.text for segment in segments).strip() + if not text: + raise RuntimeError("ASR 返回结果为空") + detected_language = getattr(info, "language", normalized_language or "unknown") + return ASRResponse(text=text, language=str(detected_language)) except HTTPException: raise - except Exception as e: # noqa: ANN001 - logger.exception("ASR 推理失败") - raise HTTPException(status_code=500, detail=f"ASR 推理失败: {e}") + except Exception as exc: + logger.exception("ASR inference failed") + raise HTTPException(status_code=500, detail=f"ASR 推理失败: {exc}") + finally: + if tmp_path and os.path.exists(tmp_path): + os.unlink(tmp_path) @generation_router.post("/tts", response_model=TTSResponse) async def tts_endpoint(req: TTSRequest): - """TTS 文字转语音端点""" return await generate_tts_response( text=req.text, instruct=req.instruct or "", @@ -479,18 +318,11 @@ async def tts_endpoint(req: TTSRequest): @generation_router.post("/asr", response_model=ASRResponse) async def asr_endpoint(req: ASRRequest): - """语音识别端点(非流式)""" audio_bytes = base64.b64decode(req.audio_base64) return await generate_asr_response(audio_bytes, req.language if req.language else None) def register_tts_asr_routes(app, include_generation_routes: bool = True): - """注册 TTS/ASR 路由到 FastAPI 应用""" app.include_router(meta_router, prefix="/v1/tts-asr") if include_generation_routes: app.include_router(generation_router, prefix="/v1/tts-asr") - - -router = APIRouter() -router.include_router(meta_router) -router.include_router(generation_router) diff --git a/src/AGENTS.md b/src/AGENTS.md index c5e0434..4e70c50 100644 --- a/src/AGENTS.md +++ b/src/AGENTS.md @@ -52,7 +52,7 @@ - 上传图片和文档 - 触发 OCR(通过 **OCRImageWrapper.vue**) - 导入导出 Markdown - - 导出 DOCX 和 PDF + - 导出 DOCX 和 PDF(纯前端,从展开后的编辑器 HTML 导出) - AI 开关 - 32 KB 大小限制 - TTS 菜单和播放器 @@ -68,7 +68,7 @@ - debounceMs - privacyMode - language - - currency + - country - backgroundType - backgroundImage - backgroundOpacity @@ -85,7 +85,7 @@ - 生成 request_id - 绑定 AbortSignal - 触发中止时的 cancel 请求 - - 读取 settings store 中的 thinking、privacy、language、currency、timezone 信息 + - 读取 settings store 中的 thinking、privacy、language、country、timezone 信息 - 向后端发 POST 并读取 JSON ## 当前真实约定 @@ -99,9 +99,12 @@ ## 容易踩坑的点 - 文档超过 32 KB 时,AI 补全会被禁用。 +- AI 开关需要同时作用到主编辑器、文档块嵌套编辑器、网页搜索块嵌套编辑器;当前通过 `llm-in-text:copilot-toggle` 广播同步。 - 在文档块、Mermaid、LaTeX 等特定上下文中,部分 AI 行为和上传行为会被禁用或改道。 - OCR 文本和文档块摘录会被注入补全上下文,但这些内容不应直接作为用户可见输出回写到文档。 -- 前端存在 /v1/export/pdf 调用,但调试前先确认后端是否真的实现了这个端点。 +- DOCX/PDF 导出当前不再依赖 `/v1/export/pdf`;导出逻辑看 `utils/richExport.js`。 +- 上传单文件限制统一为 100MB;图片上传会等待 OCR 完成并给出全局进度与失败反馈。 +- 视频上传会调用 `/v1/ocr` 的 `media_type=video` 路径,后端返回合并后的 OCR/ASR 文本,再插入文档块。 - **Web Search 块使用 ```llm-websearch fenced code 语法,触发词为 [WEBSEARCH](不区分大小写)。** - **验证码组件依赖 vue3-captcha 库,刷新和验证逻辑需保持与后端 /captcha/generate 和 /captcha/verify 接口同步。** - 当前前端多处仍保留占位 API Key 或默认值;不要把这种写法继续扩散到新代码。 diff --git a/src/components/DocBlockCrepe.vue b/src/components/DocBlockCrepe.vue index 5e9e4e2..c9af312 100644 --- a/src/components/DocBlockCrepe.vue +++ b/src/components/DocBlockCrepe.vue @@ -56,6 +56,7 @@ import { hiddenTextInputPlugin, hiddenTextNode, hiddenTextRemark, hiddenTextView import { fetchSuggestion, submitCompress, pollCompressStatus } from '../utils/api.js' import { isDocumentVisible, getRecommendedDebounce, getRecommendedSyncInterval } from '../composables/useVisibility.js' +const COPILOT_TOGGLE_EVENT = 'llm-in-text:copilot-toggle' const props = defineProps({ docType: { type: String, default: 'txt' }, docName: { type: String, default: 'document.txt' }, @@ -77,6 +78,7 @@ let crepe = null let syncTimer = null let syncingExternal = false let compressPoller = null +let copilotToggleHandler = null const handleCompress = () => { if (compressState.value !== 'idle') return @@ -233,8 +235,26 @@ onMounted(async () => { crepe.editor.action((ctx) => { const view = ctx.get(editorViewCtx) - setCopilotEnabled(view, true) + const enabled = typeof window !== 'undefined' + ? window.__LLM_IN_TEXT_COPILOT_ENABLED__ !== false + : true + setCopilotEnabled(view, enabled) + if (!enabled) { + clearGhostSuggestion(view) + } }) + + copilotToggleHandler = (event) => { + const enabled = Boolean(event?.detail?.enabled) + crepe?.editor?.action((ctx) => { + const view = ctx.get(editorViewCtx) + setCopilotEnabled(view, enabled) + if (!enabled) { + clearGhostSuggestion(view) + } + }) + } + window.addEventListener(COPILOT_TOGGLE_EVENT, copilotToggleHandler) }) onUnmounted(() => { @@ -246,6 +266,10 @@ onUnmounted(() => { compressPoller.stop() compressPoller = null } + if (copilotToggleHandler) { + window.removeEventListener(COPILOT_TOGGLE_EVENT, copilotToggleHandler) + copilotToggleHandler = null + } if (crepe) { crepe.editor.action((ctx) => { const view = ctx.get(editorViewCtx) @@ -444,4 +468,3 @@ onUnmounted(() => { display: none !important; } - diff --git a/src/components/InputBlockCrepe.vue b/src/components/InputBlockCrepe.vue new file mode 100644 index 0000000..0384459 --- /dev/null +++ b/src/components/InputBlockCrepe.vue @@ -0,0 +1,223 @@ +