5a26dfde2a
后端变更: - 新增 risk_config.py: 风险配置数据类,支持环境变量驱动 - 新增 risk_control.py: 风险控制控制器,管理并发和预算 - 新增 session_store.py: 匿名会话存储,基于 cookie 的 session ID - 新增 audit_store.py: API 审计日志存储,记录请求和 LLM 调用 - 新增 captcha_api.py: 验证码 API,用于验证用户操作真实性 - 新增 llm_policy.py: LLM 策略配置,管理 completion/pro/vision 模型 - main.py: 集成 middleware、risk/audit/session 模块 (+467/-7) - job_handlers.py: LLM 执行流程重构,新增 risk/audit 集成 (+207/-4) - llm.py: 异步客户端封装,新增 max_output_tokens 参数 (+78/-1) - job_system.py: stream_events 逻辑优化,支持心跳检测 (+12/-4) - pro_completions.py: SSE heartbeat 机制,防止连接超时 (+14/-4) - prompt.py: _normalize_preferences 支持 Mapping 类型 (+13/-0) - tts_asr.py: asyncio loop 初始化,router export (+10/-0) 前端变更: - src/components/CaptchaComponent.vue: 新增验证码组件 (NEW) - src/utils/cookie_policy.js: Cookie 策略工具 (NEW) - SettingsPanel.vue: 集成验证码组件,新增安全设置部分 (+59/-0) - MilkdownEditor.vue: 移除硬编码 API_KEY,新增 credentials (+32/-10) - ProBlockCrepe.vue: 样式简化,移除渐变动画 (+18/-4) - proBlockPlugin.ts: 重构 schema/serializer 引用方式,通过 Ctx 管理 (+40/-10) - api.js: 新增 credentials,重构 headers 条件逻辑 (+50/-14) - config.js: API 基址改为 https://api.imageteach.tech:8002 (+8/-4) - convert.js, docsApi.js, i18n.js: 新增 credentials 和验证码 i18n (+54/-12) - proAccept.js: 重构正则和转义处理,修复捕获组索引 (+14/-4) 配置和基础设施: - docker-compose.yml: 新增端口映射 8001:8001 (+2/-0) - docker/nginx.conf: 改为 307 redirect,优化代理配置 (+8/-6) - vite.config.js: 移除 proxy 配置,直接调用远程 API (+8/-4) - .env.example: 新增 VITE_API_BASE_URL, VITE_API_KEY (+3/-1) - backend/.env.example: 大量 RISK_*, SESSION_*, CORS_* 配置 (+54/-0) - pytest.ini: 扩展 coverage 范围到整个 backend,移除 fail_under (+3/-2) - .coveragerc: 移除 fail_under = 90 (+0/-1) - .gitignore: 新增 docker-data/ (+3/-0) - package.json: 新增 vue3-captcha 依赖 (+3/-1) - AGENTS.md, README.md: 更新 Docker 部署和前端网络约定 (+20/-5) - public/sw.js: Service Worker cache 版本从 v1 升级到 v2 (+0/-1) 测试变更: - test_main_endpoints.py: 新增 session/risk/audit reset,新增测试用例 (+63/-4) - test_main_cancel.py: 新增 reset 调用 (+6/-0) - test_pro_completions.py: 新增 preferences 序列化和测试 (+23/-0) 总计: 45 个文件变更,+1009/-280 行
246 lines
8.5 KiB
Python
246 lines
8.5 KiB
Python
import base64
|
|
import asyncio
|
|
import importlib
|
|
import os
|
|
import sys
|
|
from pathlib import Path
|
|
|
|
from fastapi.testclient import TestClient
|
|
|
|
os.environ["JOB_BACKEND"] = "memory"
|
|
os.environ["DOCS_BACKEND"] = "memory"
|
|
|
|
CURRENT_DIR = Path(__file__).resolve().parent
|
|
BACKEND_DIR = CURRENT_DIR.parent
|
|
if str(BACKEND_DIR) not in sys.path:
|
|
sys.path.insert(0, str(BACKEND_DIR))
|
|
|
|
import job_handlers # type: ignore
|
|
import job_system # type: ignore
|
|
import docs_store # type: ignore
|
|
import risk_control # type: ignore
|
|
import session_store # type: ignore
|
|
import audit_store # type: ignore
|
|
|
|
main = importlib.import_module("main")
|
|
|
|
HEADERS = {"X-API-Key": main.API_KEY}
|
|
|
|
|
|
def setup_function():
|
|
job_system.reset_job_manager()
|
|
docs_store.reset_document_store()
|
|
risk_control.reset_risk_controller()
|
|
session_store.reset_session_store()
|
|
audit_store.reset_audit_store()
|
|
main._handlers_registered = False
|
|
|
|
|
|
class DummyRequest:
|
|
def __init__(self, host=None, headers=None):
|
|
class Client:
|
|
pass
|
|
self.client = Client() if host is not None else None
|
|
if self.client is not None:
|
|
self.client.host = host
|
|
self.headers = headers or {}
|
|
|
|
|
|
def test_preview_short_text():
|
|
assert main._preview("Hello") == "Hello"
|
|
|
|
|
|
def test_preview_long_text_truncated():
|
|
long_text = "a" * 100
|
|
assert main._preview(long_text) == long_text[:80] + "..."
|
|
|
|
|
|
def test_sanitize_markdown_strips_image_markdown():
|
|
assert "" not in main._sanitize_converted_markdown("text ")
|
|
|
|
|
|
def test_sanitize_inline_completion_strips_prefill():
|
|
assert main.sanitize_inline_completion_content("系统非常适合写作", prefill="系统") == "非常适合写作"
|
|
|
|
|
|
def test_get_client_ip_header_overrides_host():
|
|
req = DummyRequest(host="1.2.3.4", headers={"X-Client-IP": "5.6.7.8"})
|
|
assert main.get_client_ip(req) == "5.6.7.8"
|
|
|
|
|
|
def test_post_completions_without_api_key_uses_anonymous_session(monkeypatch):
|
|
async def fake_call(*args, **kwargs):
|
|
return {"content": "系统done", "think": ""}
|
|
|
|
monkeypatch.setattr(job_handlers, "call_ollama", fake_call)
|
|
with TestClient(main.app) as client:
|
|
with client.stream("POST", "/v1/completions", json={
|
|
"prefix": "hello", "suffix": "", "languageId": "markdown",
|
|
"model_thinking": "low", "privacy_mode": True,
|
|
}) as resp:
|
|
assert resp.status_code == 200
|
|
assert main.config.session_cookie_name in resp.cookies
|
|
|
|
|
|
def test_post_completions_invalid_api_key_returns_403():
|
|
with TestClient(main.app) as client:
|
|
resp = client.post(
|
|
"/v1/completions",
|
|
headers={"X-API-Key": "invalid-key"},
|
|
json={
|
|
"prefix": "hello", "suffix": "", "languageId": "markdown",
|
|
"model_thinking": "low", "privacy_mode": True,
|
|
},
|
|
)
|
|
assert resp.status_code == 403
|
|
|
|
|
|
def test_post_completions_returns_sse_done(monkeypatch):
|
|
async def fake_call(*args, **kwargs):
|
|
return {"content": "系统done", "think": ""}
|
|
|
|
monkeypatch.setattr(job_handlers, "call_ollama", fake_call)
|
|
with TestClient(main.app) as client:
|
|
with client.stream("POST", "/v1/completions", headers=HEADERS, json={
|
|
"prefix": "hello", "suffix": "", "languageId": "markdown",
|
|
"model_thinking": "low", "privacy_mode": True,
|
|
}) as resp:
|
|
assert resp.status_code == 200
|
|
body = "".join(resp.iter_text())
|
|
assert "event: queued" in body
|
|
assert "event: started" in body
|
|
assert "event: result" in body
|
|
assert "event: done" in body
|
|
|
|
|
|
def test_stream_job_emits_keepalive_during_idle(monkeypatch):
|
|
async def fake_queue_job(*_args, **_kwargs):
|
|
return "job-keepalive"
|
|
|
|
class FakeManager:
|
|
def register_handler(self, *_args, **_kwargs):
|
|
return None
|
|
|
|
async def stream_events(self, job_id):
|
|
yield {"event": "queued", "job_id": job_id}
|
|
yield {"event": "started", "job_id": job_id}
|
|
await asyncio.sleep(0.03)
|
|
yield {"event": "done", "job_id": job_id, "result": {"content": "ok"}}
|
|
|
|
monkeypatch.setattr(main, "STREAM_HEARTBEAT_SECONDS", 0.01)
|
|
monkeypatch.setattr(main, "_queue_job", fake_queue_job)
|
|
monkeypatch.setattr(main, "get_job_manager", lambda: FakeManager())
|
|
|
|
with TestClient(main.app) as client:
|
|
with client.stream("POST", "/v1/pro/completions", headers=HEADERS, json={
|
|
"prefix": "hello", "suffix": "", "languageId": "markdown",
|
|
"instruction": "expand", "pro_thinking": "medium", "privacy_mode": True,
|
|
}) as resp:
|
|
assert resp.status_code == 200
|
|
body = "".join(resp.iter_text())
|
|
|
|
assert ": keepalive" in body
|
|
assert "event: done" in body
|
|
|
|
|
|
def test_post_ocr_mocked(monkeypatch):
|
|
async def fake_ocr(*args, **kwargs):
|
|
return "OCR result text"
|
|
|
|
monkeypatch.setattr(job_handlers, "call_vlm_ocr", fake_ocr)
|
|
img_b64 = base64.b64encode(b"pretend image data").decode()
|
|
with TestClient(main.app) as client:
|
|
with client.stream("POST", "/v1/ocr", headers=HEADERS, json={
|
|
"image": img_b64, "filename": "test.jpg", "language": "auto",
|
|
}) as resp:
|
|
assert resp.status_code == 200
|
|
body = "".join(resp.iter_text())
|
|
assert "OCR result text" in body
|
|
|
|
|
|
def test_post_convert_txt_returns_markdown():
|
|
content = base64.b64encode(b"hello world").decode()
|
|
with TestClient(main.app) as client:
|
|
with client.stream("POST", "/v1/convert", headers=HEADERS, json={
|
|
"file": content, "filename": "sample.txt",
|
|
}) as resp:
|
|
assert resp.status_code == 200
|
|
body = "".join(resp.iter_text())
|
|
assert "hello world" in body
|
|
|
|
|
|
def test_post_convert_unsupported_extension_returns_500():
|
|
content = base64.b64encode(b"data").decode()
|
|
with TestClient(main.app) as client:
|
|
resp = client.post("/v1/convert", headers=HEADERS, json={
|
|
"file": content, "filename": "sample.xlsx",
|
|
})
|
|
assert resp.status_code == 500
|
|
assert "仅支持" in resp.json()["error"]
|
|
|
|
|
|
def test_docs_nodes_crud_round_trip():
|
|
with TestClient(main.app) as client:
|
|
folder_resp = client.post("/v1/docs/folders", headers=HEADERS, json={
|
|
"name": "项目资料",
|
|
"parentId": None,
|
|
})
|
|
assert folder_resp.status_code == 200
|
|
folder = folder_resp.json()["node"]
|
|
|
|
file_resp = client.post("/v1/docs/files/text", headers=HEADERS, json={
|
|
"name": "notes.md",
|
|
"parentId": folder["id"],
|
|
"content": "# hello",
|
|
})
|
|
assert file_resp.status_code == 200
|
|
file_node = file_resp.json()["node"]
|
|
assert file_node["previewText"] == "# hello"
|
|
|
|
list_resp = client.get("/v1/docs/nodes", headers=HEADERS)
|
|
assert list_resp.status_code == 200
|
|
nodes = list_resp.json()["nodes"]
|
|
assert len(nodes) == 2
|
|
|
|
rename_resp = client.patch(f"/v1/docs/nodes/{file_node['id']}", headers=HEADERS, json={
|
|
"name": "renamed.md",
|
|
})
|
|
assert rename_resp.status_code == 200
|
|
assert rename_resp.json()["node"]["name"] == "renamed.md"
|
|
|
|
blob_resp = client.get(f"/v1/docs/files/{file_node['id']}/blob", headers=HEADERS)
|
|
assert blob_resp.status_code == 200
|
|
assert blob_resp.content == b"# hello"
|
|
|
|
delete_resp = client.delete(f"/v1/docs/nodes/{folder['id']}", headers=HEADERS)
|
|
assert delete_resp.status_code == 200
|
|
|
|
final_list = client.get("/v1/docs/nodes", headers=HEADERS)
|
|
assert final_list.status_code == 200
|
|
assert final_list.json()["nodes"] == []
|
|
|
|
|
|
def test_docs_file_upload_and_blob_replace():
|
|
with TestClient(main.app) as client:
|
|
upload_resp = client.post(
|
|
"/v1/docs/files/upload",
|
|
headers=HEADERS,
|
|
files={"file": ("image.png", b"png-bytes", "image/png")},
|
|
data={"parent_id": ""},
|
|
)
|
|
assert upload_resp.status_code == 200
|
|
node = upload_resp.json()["node"]
|
|
assert node["storageKind"] == "blob"
|
|
|
|
replace_resp = client.put(
|
|
f"/v1/docs/files/{node['id']}/blob",
|
|
headers=HEADERS,
|
|
files={"file": ("photo.jpg", b"jpeg-bytes", "image/jpeg")},
|
|
)
|
|
assert replace_resp.status_code == 200
|
|
assert replace_resp.json()["node"]["name"] == "photo.jpg"
|
|
|
|
blob_resp = client.get(f"/v1/docs/files/{node['id']}/blob", headers=HEADERS)
|
|
assert blob_resp.status_code == 200
|
|
assert blob_resp.content == b"jpeg-bytes"
|