import os import sys import types import asyncio import threading from fastapi.testclient import TestClient CURRENT_DIR = os.path.dirname(os.path.abspath(__file__)) BACKEND_DIR = os.path.abspath(os.path.join(CURRENT_DIR, "..")) if BACKEND_DIR not in sys.path: sys.path.insert(0, BACKEND_DIR) if "tts_asr" not in sys.modules: fake_tts_asr = types.ModuleType("tts_asr") fake_tts_asr.register_tts_asr_routes = lambda app: None sys.modules["tts_asr"] = fake_tts_asr import main # type: ignore import pro_completions # type: ignore HEADERS = {"X-API-Key": main.API_KEY} def _payload(): return { "prefix": "Before", "suffix": "After", "languageId": "markdown", "instruction": "expand", "pro_thinking": "medium", "privacy_mode": True, } def setup_function(): pro_completions.PRO_STATES.clear() def teardown_function(): pro_completions.PRO_STATES.clear() def test_pro_queue_full_returns_429(monkeypatch): monkeypatch.setattr(pro_completions, "PRO_QUEUE_MAX_SIZE", 0) client = TestClient(main.app) response = client.post("/v1/pro/completions", headers=HEADERS, json=_payload()) assert response.status_code == 429 assert response.json()["error"] == "PRO queue is full" def test_pro_status_missing_returns_404(): client = TestClient(main.app) response = client.get("/v1/pro/completions/status/missing", headers=HEADERS) assert response.status_code == 404 def test_pro_prompt_uses_pro_specific_instruction(): system_prompt, user_prompt = pro_completions._build_pro_prompts( prefix="欢迎使用 LLM-IN-TEXT\n\n即时可用的 LLM 系统", suffix="", language_id="markdown", instruction="", pro_thinking="high", ) combined = f"{system_prompt}\n{user_prompt}".lower() assert "[pro] model for llm-in-text" in combined assert "pro_mode: true" in combined assert "pro_thinking_level: high" in combined assert "long paragraphs or section-level output are allowed" in combined assert "highest priority" in combined assert "never copy tags to output" in combined assert "write only the markdown that belongs at the cursor" not in combined assert "continue the markdown naturally" in combined def test_pro_cancel_waits_for_stream_cleanup(monkeypatch): started = threading.Event() cleaned = threading.Event() async def fake_stream_events(*args, **kwargs): started.set() try: yield "thinking", "" while True: await asyncio.sleep(0.05) finally: cleaned.set() monkeypatch.setattr(pro_completions, "stream_ollama_events", fake_stream_events) request_id = "pro-cancel-cleanup" headers = {**HEADERS, "X-Request-Id": request_id} response_box = {} with TestClient(main.app) as client: def send_stream(): with client.stream("POST", "/v1/pro/completions", headers=headers, json=_payload()) as response: response_box["status_code"] = response.status_code response_box["body"] = "".join(response.iter_text()) stream_thread = threading.Thread(target=send_stream, daemon=True) stream_thread.start() assert started.wait(timeout=2.0) cancel_response = client.post( "/v1/pro/completions/cancel", headers=HEADERS, json={"request_id": request_id, "reason": "test"}, ) assert cancel_response.status_code == 200 assert cancel_response.json() == {"cancelled": True, "status": "ok"} assert cleaned.wait(timeout=2.0) stream_thread.join(timeout=5.0) assert not stream_thread.is_alive()