import asyncio import os import re from contextlib import suppress from typing import Any, Callable, Awaitable import markitdown from audit_store import get_audit_store from llm import call_ollama, call_vlm_ocr, stream_ollama_events from prompt import ( build_completion_prompts, build_pro_completion_prompts, prepare_prompt_context, ) from risk_config import load_risk_config from risk_control import RiskIdentity, estimate_tokens, get_risk_controller try: # pragma: no cover - optional heavy dependency path from tts_asr import generate_asr_response, generate_tts_response except Exception: # pragma: no cover generate_tts_response = None generate_asr_response = None IMAGE_MARKDOWN_RE = re.compile(r"!\[[^\]]*]\([^)]+\)") IMAGE_HTML_RE = re.compile(r"]*>", re.IGNORECASE) ALLOWED_CONVERT_EXTENSIONS = {".txt", ".docx", ".pptx", ".pdf"} _markitdown_instance = None _risk_config = load_risk_config() def _get_markitdown(): global _markitdown_instance if _markitdown_instance is None: _markitdown_instance = markitdown.MarkItDown() return _markitdown_instance def _safe_unlink(path: str | None) -> None: if not path: return with suppress(FileNotFoundError): os.unlink(path) def _sanitize_converted_markdown(text: str) -> str: value = (text or "").replace("\r\n", "\n").replace("\r", "\n") value = IMAGE_MARKDOWN_RE.sub("", value) value = IMAGE_HTML_RE.sub("", value) value = re.sub(r"\n{3,}", "\n\n", value) return value.strip() def sanitize_inline_completion_content(text: str, prefill: str = "") -> str: value = (text or "").strip() if not value: return "" fim_middle = value.rfind("<|fim_middle|>") if fim_middle >= 0: value = value[fim_middle + len("<|fim_middle|>") :] end_index = value.find("<|end|>") if end_index >= 0: value = value[:end_index] quoted = re.findall(r'"([^"]+)"', value) if quoted: value = quoted[-1] marker_index = max(value.rfind("|fim_middle|>"), value.rfind("<|start|>assistant")) if marker_index >= 0: tail = value.split(">")[-1] if tail: value = tail value = value.strip() if prefill and value.startswith(prefill): value = value[len(prefill) :] return value.strip() def _payload_identity(payload: dict[str, Any]) -> RiskIdentity: risk = payload.get("risk") or {} return RiskIdentity( request_id=risk.get("request_id") or payload["request_id"], session_hash=risk.get("session_hash", ""), ip_hash=risk.get("ip_hash", ""), route=payload.get("route", payload.get("job_type", payload.get("request_id", ""))), method="POST", ) async def _enter_llm_execution(payload: dict[str, Any], emit: Callable[[str, dict[str, Any]], Awaitable[None]]) -> tuple[RiskIdentity, dict[str, Any], list[str]]: risk = payload.get("risk") or {} identity = _payload_identity(payload) delay_ms = int(risk.get("delay_ms", 0) or 0) policy = risk.get("policy") or {} if delay_ms > 0: await emit("resource", {"phase": "delay", "delay_ms": delay_ms}) await asyncio.sleep(delay_ms / 1000.0) controller = get_risk_controller(_risk_config) lock_keys = await controller.acquire_execution_slot(identity, model=policy.get("model", "")) return identity, risk, lock_keys async def _exit_llm_execution( payload: dict[str, Any], identity: RiskIdentity, risk: dict[str, Any], lock_keys: list[str], *, status: str, actual_output_text: str = "", error_code: str = "", ) -> None: policy = (risk.get("policy") or {}) controller = get_risk_controller(_risk_config) await controller.release_execution_slot(identity, lock_keys, model=policy.get("model", "")) await controller.record_model_result(model=policy.get("model", ""), success=(status == "completed")) store = get_audit_store(os.getenv("DATABASE_URL", "").strip() or None) estimated_input_tokens = int(risk.get("estimated_input_tokens", 0) or 0) profile = policy.get("profile", "completion") pricing_out = { "completion": _risk_config.completion_output_cost_per_1k, "pro": _risk_config.pro_output_cost_per_1k, "vision": _risk_config.vision_output_cost_per_1k, }.get(profile, _risk_config.completion_output_cost_per_1k) actual_output_tokens = estimate_tokens(actual_output_text) actual_cost = round((estimated_input_tokens / 1000.0) * { "completion": _risk_config.completion_input_cost_per_1k, "pro": _risk_config.pro_input_cost_per_1k, "vision": _risk_config.vision_input_cost_per_1k, }.get(profile, _risk_config.completion_input_cost_per_1k) + (actual_output_tokens / 1000.0) * pricing_out, 8) await asyncio.to_thread( store.record_llm_call, { "request_id": payload["request_id"], "session_hash": identity.session_hash, "ip_hash": identity.ip_hash, "job_type": policy.get("job_type", ""), "model": policy.get("model", ""), "estimated_input_tokens": estimated_input_tokens, "max_output_tokens": int(policy.get("max_output_tokens", 0) or 0), "estimated_cost": float(risk.get("estimated_cost", 0.0) or 0.0), "actual_output_chars": len(actual_output_text or ""), "actual_cost": actual_cost, "status": status, "error_code": error_code, "metadata": {"profile": profile}, }, ) async def completion_handler( payload: dict[str, Any], emit: Callable[[str, dict[str, Any]], Awaitable[None]], is_cancelled: Callable[[], bool], ) -> dict[str, Any]: req = payload["request"] identity, risk, lock_keys = await _enter_llm_execution(payload, emit) system_prompt, user_prompt, prefill = build_completion_prompts( req["prefix"], req["suffix"], req.get("languageId", "markdown"), location=payload.get("location", ""), thinking_level=req.get("model_thinking", "low"), preferences=req.get("user_preferences"), ) policy = risk.get("policy") or {} try: result = await call_ollama( user_prompt, system_prompt=system_prompt, tag=f'{payload["request_id"][:8]}-completion', temperature=float(policy.get("temperature", req.get("temperature", 0.7))), thinking=policy.get("thinking"), model=policy.get("model"), prefill=prefill or None, max_output_tokens=int(policy.get("max_output_tokens", 0) or 0), ) content = sanitize_inline_completion_content(result.get("content") or "", prefill=prefill or "") if is_cancelled(): raise asyncio.CancelledError() await emit("result", {"content": content}) await _exit_llm_execution(payload, identity, risk, lock_keys, status="completed", actual_output_text=content) return {"content": content, "request_id": payload["request_id"]} except asyncio.CancelledError: await _exit_llm_execution(payload, identity, risk, lock_keys, status="cancelled", error_code="cancelled") raise except Exception: await _exit_llm_execution(payload, identity, risk, lock_keys, status="failed", error_code="llm_failed") raise async def pro_completion_handler( payload: dict[str, Any], emit: Callable[[str, dict[str, Any]], Awaitable[None]], is_cancelled: Callable[[], bool], ) -> dict[str, Any]: req = payload["request"] identity, risk, lock_keys = await _enter_llm_execution(payload, emit) system_prompt, user_prompt = build_pro_completion_prompts( prefix=req["prefix"], suffix=req["suffix"], instruction=req.get("instruction", ""), language_id=req.get("languageId", "markdown"), location=payload.get("location", ""), pro_thinking_level=req.get("pro_thinking", "medium"), preferences=req.get("user_preferences"), ) chunks: list[str] = [] policy = risk.get("policy") or {} try: async for event_type, delta in stream_ollama_events( user_prompt, system_prompt=system_prompt, tag=f'{payload["request_id"][:8]}-pro', temperature=float(policy.get("temperature", 0.7)), thinking=policy.get("thinking"), model=policy.get("model"), enable_thinking=True, max_output_tokens=int(policy.get("max_output_tokens", 0) or 0), ): if is_cancelled(): raise asyncio.CancelledError() if event_type == "thinking": await emit("progress", {"phase": "thinking"}) continue if delta: chunks.append(delta) await emit("result", {"delta": delta}) content = "".join(chunks) await _exit_llm_execution(payload, identity, risk, lock_keys, status="completed", actual_output_text=content) return {"content": content, "request_id": payload["request_id"]} except asyncio.CancelledError: await _exit_llm_execution(payload, identity, risk, lock_keys, status="cancelled", error_code="cancelled") raise except Exception: await _exit_llm_execution(payload, identity, risk, lock_keys, status="failed", error_code="llm_failed") raise async def compress_handler( payload: dict[str, Any], emit: Callable[[str, dict[str, Any]], Awaitable[None]], is_cancelled: Callable[[], bool], ) -> dict[str, Any]: content = payload["content"] doc_type = payload.get("docType", "txt") identity, risk, lock_keys = await _enter_llm_execution(payload, emit) system_prompt = ( f"你是一个专业的文档摘要助手。请将以下 {doc_type} 类型文档内容进行精简压缩," "保留核心信息和关键要点,去除冗余和啰嗦的表述。" "请直接输出压缩后的内容,不要添加任何解释性文字。" ) policy = risk.get("policy") or {} try: result = await call_ollama( content, system_prompt=system_prompt, tag=f'{payload["request_id"][:8]}-compress', model=policy.get("model"), temperature=float(policy.get("temperature", 0.2)), thinking=policy.get("thinking"), max_output_tokens=int(policy.get("max_output_tokens", 0) or 0), ) if is_cancelled(): raise asyncio.CancelledError() compressed = result.get("content") or "" await emit("result", {"content": compressed}) await _exit_llm_execution(payload, identity, risk, lock_keys, status="completed", actual_output_text=compressed) return {"content": compressed, "request_id": payload["request_id"]} except asyncio.CancelledError: await _exit_llm_execution(payload, identity, risk, lock_keys, status="cancelled", error_code="cancelled") raise except Exception: await _exit_llm_execution(payload, identity, risk, lock_keys, status="failed", error_code="llm_failed") raise async def ocr_handler( payload: dict[str, Any], emit: Callable[[str, dict[str, Any]], Awaitable[None]], is_cancelled: Callable[[], bool], ) -> dict[str, Any]: path = payload["input_path"] identity, risk, lock_keys = await _enter_llm_execution(payload, emit) try: with open(path, "rb") as handle: image_bytes = handle.read() text = await call_vlm_ocr(image_bytes, payload.get("language", "auto")) 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")} except asyncio.CancelledError: await _exit_llm_execution(payload, identity, risk, lock_keys, status="cancelled", error_code="cancelled") raise except Exception: await _exit_llm_execution(payload, identity, risk, lock_keys, status="failed", error_code="ocr_failed") raise finally: _safe_unlink(path) async def convert_handler( payload: dict[str, Any], emit: Callable[[str, dict[str, Any]], Awaitable[None]], is_cancelled: Callable[[], bool], ) -> dict[str, Any]: path = payload["input_path"] filename = payload.get("filename", "document") ext = os.path.splitext(filename)[1].lower() if ext not in ALLOWED_CONVERT_EXTENSIONS: _safe_unlink(path) raise ValueError("仅支持 txt、docx、pptx、pdf 格式") try: if ext == ".txt": with open(path, "rb") as handle: markdown = _sanitize_converted_markdown(handle.read().decode("utf-8", errors="ignore")) else: md = _get_markitdown() result = await asyncio.to_thread(md.convert, path) markdown = _sanitize_converted_markdown(result.text_content) if is_cancelled(): raise asyncio.CancelledError() await emit("result", {"markdown": markdown}) return {"markdown": markdown, "filename": filename} finally: _safe_unlink(path) async def tts_handler( payload: dict[str, Any], emit: Callable[[str, dict[str, Any]], Awaitable[None]], is_cancelled: Callable[[], bool], ) -> dict[str, Any]: if generate_tts_response is None: raise RuntimeError("TTS 功能当前不可用") response = await generate_tts_response( text=payload["text"], instruct=payload.get("instruct", ""), speaker=payload.get("speaker", "Vivian"), output_format=payload.get("format", "wav"), ) if is_cancelled(): raise asyncio.CancelledError() result = response.dict() await emit("result", result) return result async def asr_handler( payload: dict[str, Any], emit: Callable[[str, dict[str, Any]], Awaitable[None]], is_cancelled: Callable[[], bool], ) -> dict[str, Any]: if generate_asr_response is None: raise RuntimeError("ASR 功能当前不可用") path = payload["input_path"] try: with open(path, "rb") as handle: audio_bytes = handle.read() response = await generate_asr_response(audio_bytes, payload.get("language", "zh-CN")) if is_cancelled(): raise asyncio.CancelledError() result = response.dict() await emit("result", result) return result finally: _safe_unlink(path)