2026-05-31 16:38:10 +08:00
|
|
|
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
|
|
|
|
|
|
|
|
|
|
|
2026-06-02 21:23:34 +08:00
|
|
|
def test_pro_prompt_uses_pro_specific_instruction():
|
2026-05-31 16:38:10 +08:00
|
|
|
system_prompt, user_prompt = pro_completions._build_pro_prompts(
|
|
|
|
|
prefix="欢迎使用 LLM-IN-TEXT\n\n即时可用的 LLM 系统",
|
|
|
|
|
suffix="",
|
|
|
|
|
language_id="markdown",
|
|
|
|
|
instruction="",
|
2026-06-02 21:23:34 +08:00
|
|
|
pro_thinking="high",
|
2026-05-31 16:38:10 +08:00
|
|
|
)
|
|
|
|
|
combined = f"{system_prompt}\n{user_prompt}".lower()
|
2026-06-02 21:23:34 +08:00
|
|
|
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
|
2026-05-31 16:38:10 +08:00
|
|
|
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()
|