Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a75d81a7d9 | ||
|
|
5dd5a46aae | ||
|
|
1fe75e3fd2 | ||
|
|
d31cd842c5 | ||
|
|
abccb328fc | ||
|
|
ae65c64c8f | ||
|
|
3bd475dc15 | ||
|
|
78e8e3e33b | ||
|
|
9b8b10cdb1 | ||
|
|
3898530585 | ||
|
|
fcc601fcf3 | ||
|
|
2f7066aa92 | ||
|
|
7d5f4023a9 | ||
|
|
2dc984401d | ||
|
|
ed2e867db1 | ||
|
|
c6cde2500b | ||
|
|
1e32b2e0f4 | ||
|
|
ff3da5d6b1 | ||
|
|
0006e91e67 | ||
|
|
866febec21 | ||
|
|
9b50b8f0ce |
@@ -14,6 +14,12 @@ backend/.env
|
||||
# 运行期生成的 SQLite 索引(vault 下的 Markdown 测试数据需提交)
|
||||
backend/data/*.db*
|
||||
backend/data/credentials/
|
||||
# 阶段验收笔记(验收用,不提交)
|
||||
backend/data/vault/验收/
|
||||
# 本机 MCP 配置、授权状态及服务器工作目录不得提交。
|
||||
backend/data/mcp/
|
||||
server.json
|
||||
servers.json
|
||||
|
||||
# Editors and operating systems
|
||||
.idea/
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
> 本文件用于团队开发期间快速配置环境和启动项目,不是正式的项目 README。
|
||||
|
||||
> 当前基线:2026-08-30。第一阶段 Web 联调版的前端页面、Knowledge/Retrieval Core、AI/Agent Core、Extension Core、Provider 预设与本地加密凭据链路均已实现;Tauri Host、Stronghold、真实桌面文件系统和 Sync Server 尚未接入。
|
||||
> 当前基线:2026-09-03。第一阶段 Web 联调前后端已经完成;第二阶段已完成 Workspace 去 Mock、Agent Trace 持久化与 SSE 恢复、stdio MCP Bridge、隔离 Plugin Host、Plugin Command/Settings,以及独立 MCP Server 配置中心 C.1(stdio、Streamable HTTP 与旧 SSE 兼容)。真实音频、Provider 协议增强、Benchmark、导出、主题包、Trace 可视化、Mermaid 与函数图像仍在后续开发;Tauri Host、Stronghold、原生多 Vault 文件系统和 Sync Server 尚未接入。
|
||||
|
||||
## 当前目录
|
||||
|
||||
@@ -118,7 +118,7 @@ cd frontend
|
||||
pnpm test
|
||||
```
|
||||
|
||||
当前回归基线为后端 126 项测试、前端 29 项测试,且 TypeScript 类型检查和生产构建通过。测试数量会随功能增长,以本地实际输出和 CI 为准。
|
||||
当前回归基线为后端 218 项测试、前端 32 项测试,且 TypeScript 类型检查和生产构建通过。测试数量会随功能增长,以本地实际输出和 CI 为准。
|
||||
|
||||
构建产物位于 `frontend/dist`,该目录不提交到 Git。
|
||||
|
||||
@@ -134,6 +134,7 @@ pnpm test
|
||||
| [AI Core 与 Agent Core](docs/development/AI-Core与Agent-Core开发说明.md) | Provider、Agent、Tool、Permission 与 Extension Core |
|
||||
| [MCP Bridge 与 Plugin Host](docs/development/MCP-Bridge与Plugin-Host开发说明.md) | stdio MCP、隔离进程、Tool 映射、状态与错误边界 |
|
||||
| [Plugin Command 与 Settings](docs/development/Plugin-Command与Settings开发说明.md) | Command Registry、Settings Schema、Secret 引用与联调边界 |
|
||||
| [Plugin Command 与 Settings 复盘](docs/retrospectives/Plugin-Command与Settings问题与修复复盘.md) | 阶段 D 连续审阅发现的安全、事务、Schema 与运行时契约问题 |
|
||||
| [Git 使用细则](docs/guides/Git使用细则-团队开发版.md) | 分支、提交、PR、Review 与合并流程 |
|
||||
| [CI/CD 细则](docs/guides/CI-CD细则-团队开发版.md) | Gitea 流水线、质量门禁、产物、发布与回滚规则 |
|
||||
| [Agent Trace 复盘](docs/retrospectives/Agent-Core第二阶段问题与修复复盘.md) | Agent 持久化、SSE 恢复、事件契约与脱敏问题复盘 |
|
||||
@@ -147,6 +148,6 @@ pnpm test
|
||||
- 后端附件目录默认是 `backend/data/attachments`,可通过 `APP_ATTACHMENTS_PATH` 覆盖;该目录由桌面 Host 管理。
|
||||
- 跨模块接口发生变化时,需要同步更新前后端类型和 `docs` 中的接口说明。
|
||||
- 当前已实现接口见 `docs/contracts/后端接口契约-开发版.md`,第二阶段规划接口见 `docs/contracts/第二阶段接口契约-开发版.md`;已实现能力以 `/openapi.json` 为准。
|
||||
- 前端页面、交互、状态管理和第一阶段验收要求见 `docs/contracts/前端页面需求说明-开发版.md`。
|
||||
- 前端页面、交互、状态管理及当前阶段后续页面需求见 `docs/contracts/前端页面需求说明-开发版.md`。
|
||||
- 分支、提交、Pull Request、Review 和冲突处理规范见 `docs/guides/Git使用细则-团队开发版.md`。
|
||||
- CI 检查、产物、发布和回滚规范见 `docs/guides/CI-CD细则-团队开发版.md`。
|
||||
|
||||
@@ -23,7 +23,7 @@ uv run uvicorn app.main:app --reload --host 127.0.0.1 --port 8000
|
||||
uv run pytest
|
||||
```
|
||||
|
||||
当前基线为 126 项测试通过。Provider API Key 可通过前端设置页写入,也可用 `OPENAI_API_KEY`、`DEEPSEEK_API_KEY` 或 `AINOTE_CREDENTIAL_<ID>` 注入;不要把真实密钥写入仓库。`plugin.*` 是 Plugin Settings 的保留凭据命名空间,通用 Provider 凭据接口不能读写。
|
||||
当前基线为 136 项测试通过。Provider API Key 可通过前端设置页写入,也可用 `OPENAI_API_KEY`、`DEEPSEEK_API_KEY` 或 `AINOTE_CREDENTIAL_<ID>` 注入;不要把真实密钥写入仓库。`plugin.*` 是 Plugin Settings 的保留凭据命名空间,通用 Provider 凭据接口不能读写。
|
||||
|
||||
团队接口清单见 `../docs/contracts/后端接口契约-开发版.md`,机器可读契约以运行时的 `/openapi.json` 为准。
|
||||
|
||||
|
||||
@@ -160,10 +160,11 @@ def read_attachment(arguments: AttachmentReadArguments, _: ToolExecutionContext)
|
||||
return attachment_service.read_attachment(**arguments.model_dump())
|
||||
|
||||
|
||||
def transcribe_audio(arguments: AudioTranscribeArguments, _: ToolExecutionContext) -> dict:
|
||||
return transcription_service.create_transcription(
|
||||
async def transcribe_audio(arguments: AudioTranscribeArguments, _: ToolExecutionContext) -> dict:
|
||||
job = await transcription_service.create_transcription(
|
||||
arguments.attachment_id, arguments.language
|
||||
).model_dump(mode="json")
|
||||
)
|
||||
return job.model_dump(mode="json")
|
||||
|
||||
|
||||
def _register(
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
"""Benchmark 服务:RAG / Agent 数据集注册、指标计算与运行管理。
|
||||
|
||||
模块划分:
|
||||
- metrics.py 纯函数指标(Hit@K / Recall@K / MRR / CitationHit / 分位数)
|
||||
- datasets.py 受控目录的 Dataset 注册与校验
|
||||
- rag.py RAG Benchmark Runner(调用 retrieval.engine.search)
|
||||
- service.py 运行注册表、配置快照与报告组装
|
||||
"""
|
||||
@@ -0,0 +1,198 @@
|
||||
"""Benchmark Dataset 注册:从受控目录加载 JSON 数据集并校验。
|
||||
|
||||
Dataset 只能来自配置目录(settings.benchmark_datasets_path),API 不接受调用方提交
|
||||
任意文件路径。目录不存在或为空时按「无数据集」处理,不报错。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
|
||||
from pydantic import BaseModel, Field, ValidationError
|
||||
|
||||
from app.config import get_settings
|
||||
from app.contracts import (
|
||||
BenchmarkDatasetInfo,
|
||||
BenchmarkKind,
|
||||
RAGDatasetCase,
|
||||
)
|
||||
from app.errors import ApiError
|
||||
|
||||
|
||||
@dataclass
|
||||
class RAGDataset:
|
||||
"""内存中的 RAG 数据集:元信息 + 已校验的 Case 列表 + 内容哈希。"""
|
||||
|
||||
dataset_id: str
|
||||
kind: BenchmarkKind
|
||||
version: str
|
||||
description: str
|
||||
cases: list[RAGDatasetCase] = field(default_factory=list)
|
||||
content_hash: str = ""
|
||||
|
||||
|
||||
class _DatasetMeta(BaseModel):
|
||||
"""Dataset 元数据的最小校验模型。
|
||||
|
||||
list_datasets 用它逐文件校验元信息字段结构,把「合法 JSON 但字段类型错误」
|
||||
(如 cases: 42)这类损坏文件隔离掉,而不是让 len() 抛 TypeError 拖垮整个列表。
|
||||
"""
|
||||
|
||||
dataset_id: str = Field(min_length=1)
|
||||
kind: str = ""
|
||||
version: str = ""
|
||||
description: str = ""
|
||||
cases: list = Field(default_factory=list)
|
||||
|
||||
|
||||
def _datasets_dir() -> Path:
|
||||
return get_settings().benchmark_datasets_path
|
||||
|
||||
|
||||
def _dataset_files() -> list[Path]:
|
||||
directory = _datasets_dir()
|
||||
if not directory.is_dir():
|
||||
return []
|
||||
return sorted(directory.glob("*.json"))
|
||||
|
||||
|
||||
def _content_hash(raw: bytes) -> str:
|
||||
return "sha256:" + hashlib.sha256(raw).hexdigest()
|
||||
|
||||
|
||||
def _read_json(path: Path) -> tuple[dict, bytes]:
|
||||
"""读取并解析 JSON 文件,返回 (dict, 原始字节);非法 JSON 抛 BENCHMARK_DATASET_INVALID。"""
|
||||
try:
|
||||
raw_bytes = path.read_bytes()
|
||||
return json.loads(raw_bytes.decode("utf-8")), raw_bytes
|
||||
except (json.JSONDecodeError, OSError, UnicodeDecodeError) as exc:
|
||||
raise ApiError(
|
||||
422,
|
||||
"BENCHMARK_DATASET_INVALID",
|
||||
f"Dataset file is not valid JSON: {path.name}",
|
||||
{"path": str(path)},
|
||||
) from exc
|
||||
|
||||
|
||||
def _dataset_from_raw(raw: dict, raw_bytes: bytes, kind: BenchmarkKind) -> RAGDataset:
|
||||
"""把单个数据集 JSON 解析为 RAGDataset,非法结构抛 BENCHMARK_DATASET_INVALID。"""
|
||||
dataset_id = raw.get("dataset_id")
|
||||
if not isinstance(dataset_id, str) or not dataset_id:
|
||||
raise ApiError(
|
||||
422,
|
||||
"BENCHMARK_DATASET_INVALID",
|
||||
"Dataset must declare a non-empty string 'dataset_id'.",
|
||||
{},
|
||||
)
|
||||
file_kind = raw.get("kind", kind.value)
|
||||
if file_kind != kind.value:
|
||||
raise ApiError(
|
||||
422,
|
||||
"BENCHMARK_DATASET_INVALID",
|
||||
f"Dataset kind mismatch: expected '{kind.value}', got '{file_kind}'.",
|
||||
{"dataset_id": dataset_id},
|
||||
)
|
||||
raw_cases = raw.get("cases")
|
||||
if not isinstance(raw_cases, list) or not raw_cases:
|
||||
raise ApiError(
|
||||
422,
|
||||
"BENCHMARK_DATASET_INVALID",
|
||||
"Dataset 'cases' must be a non-empty list.",
|
||||
{"dataset_id": dataset_id},
|
||||
)
|
||||
|
||||
cases: list[RAGDatasetCase] = []
|
||||
for index, case in enumerate(raw_cases):
|
||||
try:
|
||||
parsed = RAGDatasetCase.model_validate(case)
|
||||
except ValidationError as exc:
|
||||
raise ApiError(
|
||||
422,
|
||||
"BENCHMARK_DATASET_INVALID",
|
||||
f"Dataset case #{index} is invalid.",
|
||||
{"dataset_id": dataset_id, "case_index": index, "errors": exc.errors()},
|
||||
) from exc
|
||||
# 每个 Case 至少要声明一个期望 id,否则无法计算命中/召回
|
||||
if not parsed.expected_note_ids and not parsed.expected_block_ids:
|
||||
raise ApiError(
|
||||
422,
|
||||
"BENCHMARK_DATASET_INVALID",
|
||||
f"Dataset case '{parsed.case_id}' must declare expected_note_ids or expected_block_ids.",
|
||||
{"dataset_id": dataset_id, "case_id": parsed.case_id},
|
||||
)
|
||||
# citation_required=true 时必须声明 expected_block_ids,否则无法计算 Citation Hit Rate
|
||||
if parsed.citation_required and not parsed.expected_block_ids:
|
||||
raise ApiError(
|
||||
422,
|
||||
"BENCHMARK_DATASET_INVALID",
|
||||
f"Dataset case '{parsed.case_id}' requires expected_block_ids when citation_required is true.",
|
||||
{"dataset_id": dataset_id, "case_id": parsed.case_id},
|
||||
)
|
||||
cases.append(parsed)
|
||||
|
||||
return RAGDataset(
|
||||
dataset_id=dataset_id,
|
||||
kind=kind,
|
||||
version=str(raw.get("version", "")),
|
||||
description=str(raw.get("description", "")),
|
||||
cases=cases,
|
||||
content_hash=_content_hash(raw_bytes),
|
||||
)
|
||||
|
||||
|
||||
def list_datasets(kind: BenchmarkKind) -> list[BenchmarkDatasetInfo]:
|
||||
"""枚举受控目录下指定 kind 的数据集元信息(不含 Case 内容)。
|
||||
|
||||
逐文件用 _DatasetMeta 校验元信息字段结构,单个损坏文件隔离跳过而非整体失败,
|
||||
保证列表接口健壮;损坏细节由 load_dataset 抛出。
|
||||
"""
|
||||
infos: list[BenchmarkDatasetInfo] = []
|
||||
for path in _dataset_files():
|
||||
try:
|
||||
raw, raw_bytes = _read_json(path)
|
||||
meta = _DatasetMeta.model_validate(raw)
|
||||
except (ApiError, ValidationError):
|
||||
continue
|
||||
if meta.kind not in ("", kind.value):
|
||||
continue
|
||||
infos.append(
|
||||
BenchmarkDatasetInfo(
|
||||
dataset_id=meta.dataset_id,
|
||||
kind=kind,
|
||||
version=meta.version,
|
||||
description=meta.description,
|
||||
case_count=len(meta.cases),
|
||||
content_hash=_content_hash(raw_bytes),
|
||||
)
|
||||
)
|
||||
return infos
|
||||
|
||||
|
||||
def load_dataset(dataset_id: str, kind: BenchmarkKind) -> RAGDataset:
|
||||
"""按文件名加载并校验数据集;找不到抛 BENCHMARK_DATASET_NOT_FOUND。
|
||||
|
||||
只读取与请求 dataset_id 同名的文件({dataset_id}.json),无关文件的损坏(JSON 语法
|
||||
错误、UTF-8 解码错误、顶层非对象)不会阻断目标数据集加载;只有目标文件本身损坏
|
||||
才抛 BENCHMARK_DATASET_INVALID。按现有文件 stem 精确匹配,不拼接调用方传入的路径。
|
||||
"""
|
||||
for path in _dataset_files():
|
||||
if path.stem != dataset_id:
|
||||
continue
|
||||
raw, raw_bytes = _read_json(path)
|
||||
if not isinstance(raw, dict):
|
||||
raise ApiError(
|
||||
422,
|
||||
"BENCHMARK_DATASET_INVALID",
|
||||
"Dataset top-level must be a JSON object.",
|
||||
{"dataset_id": dataset_id, "path": path.name},
|
||||
)
|
||||
return _dataset_from_raw(raw, raw_bytes, kind)
|
||||
raise ApiError(
|
||||
404,
|
||||
"BENCHMARK_DATASET_NOT_FOUND",
|
||||
f"Benchmark dataset does not exist: {dataset_id}",
|
||||
{"dataset_id": dataset_id, "kind": kind.value},
|
||||
)
|
||||
@@ -0,0 +1,58 @@
|
||||
"""Benchmark 指标纯函数。
|
||||
|
||||
所有指标只依赖「按相关性降序的 retrieved id 列表」和「期望 id 集合」,不接触任何
|
||||
外部状态,便于单元测试与未来 Agent Benchmark 复用。retrieved 顺序越靠前越相关。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
def hit_at_k(retrieved: list[str], expected: set[str], k: int) -> bool:
|
||||
"""前 k 个结果里是否命中任意期望 id(用于 Hit@1 / Hit@5)。"""
|
||||
return any(item in expected for item in retrieved[:k])
|
||||
|
||||
|
||||
def recall_at_k(retrieved: list[str], expected: set[str], k: int) -> float:
|
||||
"""前 k 个结果召回的期望 id 占比;期望为空时视为 0。
|
||||
|
||||
结果先去重:检索结果是 Block 级,同一 Note 可能经多个 Block 重复出现,
|
||||
直接逐项计数会把同一 Note 算多次、导致 Recall 超过 1。
|
||||
"""
|
||||
if not expected:
|
||||
return 0.0
|
||||
return len(set(retrieved[:k]) & expected) / len(expected)
|
||||
|
||||
|
||||
def reciprocal_rank(retrieved: list[str], expected: set[str]) -> float:
|
||||
"""首个命中的倒数排名;未命中返回 0。rank 从 1 开始。"""
|
||||
for rank, item in enumerate(retrieved, start=1):
|
||||
if item in expected:
|
||||
return 1.0 / rank
|
||||
return 0.0
|
||||
|
||||
|
||||
def citation_hit(retrieved_block_ids: list[str], expected: set[str]) -> bool:
|
||||
"""首条结果的 block_id 是否为期望引用块(Citation Hit Rate 的逐 Case 判据)。"""
|
||||
if not retrieved_block_ids or not expected:
|
||||
return False
|
||||
return retrieved_block_ids[0] in expected
|
||||
|
||||
|
||||
def mean(values: list[float]) -> float:
|
||||
return sum(values) / len(values) if values else 0.0
|
||||
|
||||
|
||||
def percentile(values: list[float], p: float) -> float:
|
||||
"""线性插值分位数(p ∈ [0, 100]),用于 P50 / P95 延迟。空列表返回 0。"""
|
||||
if not values:
|
||||
return 0.0
|
||||
ordered = sorted(values)
|
||||
if len(ordered) == 1:
|
||||
return ordered[0]
|
||||
rank = (len(ordered) - 1) * (p / 100.0)
|
||||
lo = int(rank)
|
||||
hi = lo + 1
|
||||
if hi >= len(ordered):
|
||||
return ordered[-1]
|
||||
frac = rank - lo
|
||||
return ordered[lo] + (ordered[hi] - ordered[lo]) * frac
|
||||
@@ -0,0 +1,163 @@
|
||||
"""RAG Benchmark Runner:调用检索引擎对数据集逐 Case 求值并聚合指标。
|
||||
|
||||
只读操作,直接复用 app.retrieval.engine 的 search(),不旁路检索链路。指标按
|
||||
(mode, case, repeat) 逐样本计算,再按 mode 聚合;失败样本按零分计入质量指标分母,
|
||||
避免把执行失败误判为检索质量(同时保留 total/successful/failed/failure_rate)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
|
||||
from app import repository
|
||||
from app.benchmarks import metrics as m
|
||||
from app.benchmarks.datasets import RAGDataset
|
||||
from app.contracts import (
|
||||
RAGCaseResult,
|
||||
RAGDatasetCase,
|
||||
RAGMetrics,
|
||||
RAGRunRequest,
|
||||
SearchMode,
|
||||
SearchRequest,
|
||||
)
|
||||
from app.retrieval.engine import engine
|
||||
from app.retrieval.provenance import capture_embedding
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class BenchmarkCancelled(Exception):
|
||||
"""运行在 Case 之间被取消时抛出,用于中断后台执行并标记 cancelled。"""
|
||||
|
||||
|
||||
async def run_rag(
|
||||
dataset: RAGDataset,
|
||||
request: RAGRunRequest,
|
||||
on_case: Callable[[RAGCaseResult, int, int], None] | None = None,
|
||||
should_cancel: Callable[[], bool] | None = None,
|
||||
) -> tuple[dict[str, RAGMetrics], list[RAGCaseResult]]:
|
||||
"""执行 RAG Benchmark,返回 (按 mode 聚合的指标, 全部逐样本结果)。
|
||||
|
||||
on_case 在每个样本求值完成后回调 (result, done, total),供上层更新进度与事件。
|
||||
should_cancel 在每个样本开始前被检查;返回 True 时抛出 BenchmarkCancelled 中断运行。
|
||||
"""
|
||||
total = len(request.modes) * len(dataset.cases) * request.repeat
|
||||
done = 0
|
||||
results: list[RAGCaseResult] = []
|
||||
|
||||
for mode in request.modes:
|
||||
for case in dataset.cases:
|
||||
expected_notes = _expected_notes(case)
|
||||
for repeat in range(request.repeat):
|
||||
# 让出事件循环:使运行中取消、SSE 进度与并发 API 请求能及时得到调度
|
||||
await asyncio.sleep(0)
|
||||
if should_cancel is not None and should_cancel():
|
||||
raise BenchmarkCancelled()
|
||||
result = await _evaluate_one(case, mode, request, repeat, expected_notes)
|
||||
results.append(result)
|
||||
done += 1
|
||||
if on_case is not None:
|
||||
on_case(result, done, total)
|
||||
|
||||
metrics_by_mode = {mode.value: _aggregate(results, mode) for mode in request.modes}
|
||||
return metrics_by_mode, results
|
||||
|
||||
|
||||
def _expected_notes(case: RAGDatasetCase) -> set[str]:
|
||||
"""返回笔记级期望 id;仅标注块 ID 时从块反查所属笔记,避免把标注缺失误判为检索失败。"""
|
||||
if case.expected_note_ids:
|
||||
return set(case.expected_note_ids)
|
||||
return {hit.note_id for hit in repository.get_block_hits(case.expected_block_ids)}
|
||||
|
||||
|
||||
async def _evaluate_one(
|
||||
case: RAGDatasetCase,
|
||||
mode: SearchMode,
|
||||
request: RAGRunRequest,
|
||||
repeat: int,
|
||||
expected_notes: set[str],
|
||||
) -> RAGCaseResult:
|
||||
search_request = SearchRequest(
|
||||
query=case.query,
|
||||
mode=mode,
|
||||
limit=request.retrieval.top_k,
|
||||
include_snippet=False,
|
||||
rrf_k=request.retrieval.rrf_k,
|
||||
rerank=request.retrieval.rerank,
|
||||
rerank_candidates=request.retrieval.rerank_candidates,
|
||||
score_threshold=request.retrieval.score_threshold,
|
||||
)
|
||||
start = time.perf_counter()
|
||||
embedding = {}
|
||||
try:
|
||||
with capture_embedding() as embedding:
|
||||
response = await engine.search(search_request)
|
||||
latency_ms = (time.perf_counter() - start) * 1000.0
|
||||
except Exception as exc: # 单个样本失败不中断整个 Benchmark
|
||||
# 详细异常只进日志,公开响应只带项目错误码与安全消息,避免泄露路径/SQL 等敏感信息
|
||||
logger.warning(
|
||||
"RAG case evaluation failed: case=%s mode=%s", case.case_id, mode.value,
|
||||
exc_info=exc,
|
||||
)
|
||||
return RAGCaseResult(
|
||||
embedding=embedding,
|
||||
case_id=case.case_id,
|
||||
mode=mode,
|
||||
repeat=repeat,
|
||||
latency_ms=(time.perf_counter() - start) * 1000.0,
|
||||
citation_applicable=case.citation_required,
|
||||
error="RAG case evaluation failed.",
|
||||
error_code="BENCHMARK_CASE_EVALUATION_FAILED",
|
||||
)
|
||||
|
||||
retrieved_note_ids = [item.note_id for item in response.items]
|
||||
retrieved_block_ids = [item.block_id for item in response.items]
|
||||
expected_blocks = set(case.expected_block_ids)
|
||||
k = request.retrieval.top_k
|
||||
|
||||
return RAGCaseResult(
|
||||
embedding=embedding,
|
||||
case_id=case.case_id,
|
||||
mode=mode,
|
||||
repeat=repeat,
|
||||
latency_ms=latency_ms,
|
||||
retrieved_note_ids=retrieved_note_ids,
|
||||
retrieved_block_ids=retrieved_block_ids,
|
||||
hit_at_1=m.hit_at_k(retrieved_note_ids, expected_notes, 1),
|
||||
hit_at_5=m.hit_at_k(retrieved_note_ids, expected_notes, 5),
|
||||
recall=m.recall_at_k(retrieved_note_ids, expected_notes, k),
|
||||
reciprocal_rank=m.reciprocal_rank(retrieved_note_ids, expected_notes),
|
||||
citation_hit=m.citation_hit(retrieved_block_ids, expected_blocks),
|
||||
citation_applicable=case.citation_required,
|
||||
)
|
||||
|
||||
|
||||
def _aggregate(cases: list[RAGCaseResult], mode: SearchMode) -> RAGMetrics:
|
||||
samples = [c for c in cases if c.mode == mode]
|
||||
total = len(samples)
|
||||
failed = sum(1 for c in samples if c.error is not None)
|
||||
successful = total - failed
|
||||
if total == 0:
|
||||
return RAGMetrics()
|
||||
|
||||
# 延迟只统计成功样本;失败样本按零分计入质量指标分母,避免汇总虚高
|
||||
latencies = [c.latency_ms for c in samples if c.error is None]
|
||||
citation_samples = [c for c in samples if c.citation_applicable]
|
||||
return RAGMetrics(
|
||||
hit_at_1=m.mean([1.0 if (c.error is None and c.hit_at_1) else 0.0 for c in samples]),
|
||||
hit_at_5=m.mean([1.0 if (c.error is None and c.hit_at_5) else 0.0 for c in samples]),
|
||||
recall_at_k=m.mean([c.recall if c.error is None else 0.0 for c in samples]),
|
||||
mrr=m.mean([c.reciprocal_rank if c.error is None else 0.0 for c in samples]),
|
||||
citation_hit_rate=m.mean(
|
||||
[1.0 if (c.error is None and c.citation_hit) else 0.0 for c in citation_samples]
|
||||
),
|
||||
p50_latency_ms=m.percentile(latencies, 50.0),
|
||||
p95_latency_ms=m.percentile(latencies, 95.0),
|
||||
total_cases=total,
|
||||
successful_cases=successful,
|
||||
failed_cases=failed,
|
||||
failure_rate=failed / total,
|
||||
)
|
||||
@@ -0,0 +1,349 @@
|
||||
"""Benchmark 服务:运行注册表、配置快照与报告组装。
|
||||
|
||||
RAG Benchmark 采用「创建即返回 queued、后台 Task 异步执行」的模式(与 index_service
|
||||
的 rebuild 一致):POST 创建后立即返回 202 queued 的 BenchmarkRun,由受管 asyncio.Task
|
||||
在后台逐 Case 求值,进度与事件实时写入内存注册表,供 SSE 订阅。运行记录、事件与报告
|
||||
暂存内存(_runs/_events/_reports),不持久化到 SQLite;后续接入异步任务队列时再落库。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import sys
|
||||
from datetime import datetime, timezone
|
||||
from uuid import uuid4
|
||||
|
||||
from app import repository
|
||||
from app.benchmarks import datasets
|
||||
from app.benchmarks.datasets import RAGDataset
|
||||
from app.benchmarks.rag import BenchmarkCancelled, run_rag
|
||||
from app.config import get_settings
|
||||
from app.contracts import (
|
||||
BenchmarkEvent,
|
||||
BenchmarkEventType,
|
||||
BenchmarkKind,
|
||||
BenchmarkReport,
|
||||
BenchmarkRun,
|
||||
BenchmarkStatus,
|
||||
RAGCaseResult,
|
||||
RAGMetrics,
|
||||
RAGRunRequest,
|
||||
SearchMode,
|
||||
)
|
||||
from app.errors import ApiError
|
||||
from app.retrieval.engine import engine
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_runs: dict[str, BenchmarkRun] = {}
|
||||
_events: dict[str, list[BenchmarkEvent]] = {}
|
||||
_reports: dict[str, BenchmarkReport] = {}
|
||||
_tasks: dict[str, asyncio.Task] = {}
|
||||
_subscribers: dict[str, list[asyncio.Queue[BenchmarkEvent]]] = {}
|
||||
_cancel_flags: dict[str, asyncio.Event] = {}
|
||||
MAX_RUNS = 100
|
||||
|
||||
|
||||
def _now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
def _forget(run_id: str) -> None:
|
||||
"""移除一条 run 的全部内存态;仅在 run 处于终态时调用,避免打断活动任务。"""
|
||||
_runs.pop(run_id, None)
|
||||
_events.pop(run_id, None)
|
||||
_reports.pop(run_id, None)
|
||||
_tasks.pop(run_id, None)
|
||||
_subscribers.pop(run_id, None)
|
||||
_cancel_flags.pop(run_id, None)
|
||||
|
||||
|
||||
def _evict_terminal() -> bool:
|
||||
"""超过容量时淘汰最旧的终态 run;全部为活动 run 无法淘汰时返回 False。
|
||||
|
||||
绝不能删除仍在运行(queued/running)的 run:那会连带移除其 _cancel_flags 与
|
||||
_subscribers,使后台 Task 访问时抛出 KeyError。
|
||||
"""
|
||||
terminal = (BenchmarkStatus.completed, BenchmarkStatus.failed, BenchmarkStatus.cancelled)
|
||||
while len(_runs) >= MAX_RUNS:
|
||||
victim = next(
|
||||
(rid for rid, run in _runs.items() if run.status in terminal), None
|
||||
)
|
||||
if victim is None:
|
||||
return False
|
||||
_forget(victim)
|
||||
return True
|
||||
|
||||
|
||||
def _config_snapshot(request: RAGRunRequest, dataset: RAGDataset) -> dict:
|
||||
"""记录运行时的模型 / 索引 / 环境信息,保证报告可解释、可复现。"""
|
||||
settings = get_settings()
|
||||
return {
|
||||
"dataset_id": dataset.dataset_id,
|
||||
"dataset_hash": dataset.content_hash,
|
||||
"dataset_version": dataset.version,
|
||||
"modes": [m.value for m in request.modes],
|
||||
"retrieval": request.retrieval.model_dump(),
|
||||
"repeat": request.repeat,
|
||||
"embedding": {"policy": "per_case", "details": "cases[].embedding"},
|
||||
"local_embedding": {
|
||||
"model_id": engine.embedding.model_id,
|
||||
"version": engine.embedding.version,
|
||||
"dim": engine.embedding.dim,
|
||||
},
|
||||
"reranker": {
|
||||
"model_id": engine.reranker.model_id,
|
||||
"version": engine.reranker.version,
|
||||
},
|
||||
"index_meta": repository.get_index_meta(),
|
||||
"app": {"version": settings.version, "environment": settings.environment},
|
||||
"python": sys.version.split()[0],
|
||||
"metadata": request.metadata,
|
||||
}
|
||||
|
||||
|
||||
async def _validate_index_compatibility(request: RAGRunRequest) -> None:
|
||||
"""创建 RAG Run 前校验索引已建立且与当前 Embedding 模型/维度兼容。
|
||||
|
||||
空索引或不兼容索引会让所有模式得到全 0 指标,把环境/索引错误误判为检索质量差,
|
||||
故在创建时即拒绝,返回 BENCHMARK_INDEX_INCOMPATIBLE。
|
||||
"""
|
||||
stats = repository.stats()
|
||||
meta = repository.get_index_meta()
|
||||
needs_vector = any(m in (SearchMode.vector, SearchMode.hybrid) for m in request.modes)
|
||||
|
||||
reasons: list[str] = []
|
||||
if stats["blocks"] == 0:
|
||||
reasons.append("index is empty (no indexed blocks; run /api/index/rebuild first)")
|
||||
if needs_vector:
|
||||
if meta.get("embedding_model") != engine.embedding.model_id:
|
||||
reasons.append(
|
||||
f"embedding model mismatch: index={meta.get('embedding_model')!r}, "
|
||||
f"engine={engine.embedding.model_id!r}"
|
||||
)
|
||||
if meta.get("embedding_dim") != str(engine.embedding.dim):
|
||||
reasons.append(
|
||||
f"embedding dimension mismatch: index={meta.get('embedding_dim')!r}, "
|
||||
f"engine={engine.embedding.dim}"
|
||||
)
|
||||
if await engine.vector_store.count() == 0:
|
||||
reasons.append("vector index is empty")
|
||||
if reasons:
|
||||
raise ApiError(
|
||||
409,
|
||||
"BENCHMARK_INDEX_INCOMPATIBLE",
|
||||
"Benchmark index is not built or is incompatible with the current retrieval engine.",
|
||||
{"reasons": reasons},
|
||||
)
|
||||
|
||||
|
||||
async def create_rag_run(request: RAGRunRequest) -> BenchmarkRun:
|
||||
"""创建一次 RAG Benchmark,立即返回 queued 的 BenchmarkRun,由后台 Task 执行。"""
|
||||
dataset = datasets.load_dataset(request.dataset_id, BenchmarkKind.rag)
|
||||
await _validate_index_compatibility(request)
|
||||
|
||||
# 容量检查:先淘汰终态 run 腾空间;满容量且全为活动 run 时拒绝创建
|
||||
if not _evict_terminal():
|
||||
raise ApiError(
|
||||
429,
|
||||
"BENCHMARK_CAPACITY_EXCEEDED",
|
||||
"Benchmark run capacity exceeded; wait for active runs to finish.",
|
||||
{},
|
||||
)
|
||||
|
||||
run_id = "benchmark_" + uuid4().hex[:12]
|
||||
snapshot = _config_snapshot(request, dataset)
|
||||
run = BenchmarkRun(
|
||||
run_id=run_id,
|
||||
kind=BenchmarkKind.rag,
|
||||
dataset_id=dataset.dataset_id,
|
||||
dataset_hash=dataset.content_hash,
|
||||
status=BenchmarkStatus.queued,
|
||||
progress=0.0,
|
||||
config_snapshot=snapshot,
|
||||
created_at=_now(),
|
||||
)
|
||||
_runs[run_id] = run
|
||||
_events[run_id] = []
|
||||
_subscribers[run_id] = []
|
||||
_cancel_flags[run_id] = asyncio.Event()
|
||||
_tasks[run_id] = asyncio.create_task(_execute_rag(run_id, request, dataset, snapshot))
|
||||
return run
|
||||
|
||||
|
||||
async def _execute_rag(
|
||||
run_id: str, request: RAGRunRequest, dataset: RAGDataset, snapshot: dict
|
||||
) -> None:
|
||||
"""后台执行 RAG Benchmark,实时更新进度/事件,结束后写入报告并关闭订阅。"""
|
||||
cancel_event = _cancel_flags[run_id]
|
||||
|
||||
def emit(event_type: BenchmarkEventType, data: dict) -> None:
|
||||
sequence = len(_events[run_id])
|
||||
event = BenchmarkEvent(
|
||||
event=event_type, run_id=run_id, sequence=sequence, data=data, timestamp=_now()
|
||||
)
|
||||
_events[run_id].append(event)
|
||||
for queue in _subscribers.get(run_id, []):
|
||||
queue.put_nowait(event)
|
||||
|
||||
def finish() -> None:
|
||||
_subscribers.pop(run_id, None)
|
||||
_cancel_flags.pop(run_id, None)
|
||||
|
||||
_runs[run_id] = _runs[run_id].model_copy(
|
||||
update={"status": BenchmarkStatus.running, "started_at": _now()}
|
||||
)
|
||||
emit(
|
||||
BenchmarkEventType.run_started,
|
||||
{"dataset_id": dataset.dataset_id, "modes": [m.value for m in request.modes]},
|
||||
)
|
||||
total = len(request.modes) * len(dataset.cases) * request.repeat
|
||||
|
||||
def on_case(result: RAGCaseResult, done: int, _total: int) -> None:
|
||||
progress = done / total if total else 1.0
|
||||
_runs[run_id] = _runs[run_id].model_copy(update={"progress": progress})
|
||||
emit(BenchmarkEventType.case_completed, result.model_dump(mode="json"))
|
||||
|
||||
try:
|
||||
metrics_by_mode, results = await run_rag(
|
||||
dataset,
|
||||
request,
|
||||
on_case=on_case,
|
||||
should_cancel=cancel_event.is_set,
|
||||
)
|
||||
except BenchmarkCancelled:
|
||||
_runs[run_id] = _runs[run_id].model_copy(
|
||||
update={
|
||||
"status": BenchmarkStatus.cancelled,
|
||||
"progress": 1.0,
|
||||
"completed_at": _now(),
|
||||
}
|
||||
)
|
||||
emit(BenchmarkEventType.run_cancelled, {"status": BenchmarkStatus.cancelled.value})
|
||||
_reports[run_id] = BenchmarkReport(
|
||||
run_id=run_id,
|
||||
kind=BenchmarkKind.rag,
|
||||
dataset_id=dataset.dataset_id,
|
||||
dataset_hash=dataset.content_hash,
|
||||
status=BenchmarkStatus.cancelled,
|
||||
config_snapshot=snapshot,
|
||||
)
|
||||
finish()
|
||||
return
|
||||
except Exception as exc: # 单次运行失败不拖垮服务,记录错误后结束
|
||||
# 详细异常只进日志,公开响应仅带项目错误码与安全消息,避免泄露路径/SQL 等敏感信息
|
||||
logger.exception("Benchmark run failed: run_id=%s", run_id)
|
||||
_runs[run_id] = _runs[run_id].model_copy(
|
||||
update={
|
||||
"status": BenchmarkStatus.failed,
|
||||
"progress": 1.0,
|
||||
"error": "Benchmark run failed.",
|
||||
"error_code": "BENCHMARK_RUN_FAILED",
|
||||
"completed_at": _now(),
|
||||
}
|
||||
)
|
||||
emit(
|
||||
BenchmarkEventType.run_failed,
|
||||
{"error": "Benchmark run failed.", "error_code": "BENCHMARK_RUN_FAILED"},
|
||||
)
|
||||
_reports[run_id] = BenchmarkReport(
|
||||
run_id=run_id,
|
||||
kind=BenchmarkKind.rag,
|
||||
dataset_id=dataset.dataset_id,
|
||||
dataset_hash=dataset.content_hash,
|
||||
status=BenchmarkStatus.failed,
|
||||
config_snapshot=snapshot,
|
||||
error="Benchmark run failed.",
|
||||
error_code="BENCHMARK_RUN_FAILED",
|
||||
)
|
||||
finish()
|
||||
return
|
||||
|
||||
metrics = {mode: m.model_dump() for mode, m in metrics_by_mode.items()}
|
||||
_runs[run_id] = _runs[run_id].model_copy(
|
||||
update={
|
||||
"status": BenchmarkStatus.completed,
|
||||
"progress": 1.0,
|
||||
"metrics": metrics,
|
||||
"completed_at": _now(),
|
||||
}
|
||||
)
|
||||
emit(BenchmarkEventType.run_completed, {"metrics": metrics})
|
||||
_reports[run_id] = BenchmarkReport(
|
||||
run_id=run_id,
|
||||
kind=BenchmarkKind.rag,
|
||||
dataset_id=dataset.dataset_id,
|
||||
dataset_hash=dataset.content_hash,
|
||||
status=BenchmarkStatus.completed,
|
||||
config_snapshot=snapshot,
|
||||
metrics=metrics,
|
||||
cases=results,
|
||||
)
|
||||
finish()
|
||||
|
||||
|
||||
def list_runs(
|
||||
kind: BenchmarkKind | None = None,
|
||||
status: BenchmarkStatus | None = None,
|
||||
limit: int = 50,
|
||||
offset: int = 0,
|
||||
) -> tuple[list[BenchmarkRun], int]:
|
||||
runs = list(_runs.values())
|
||||
if kind is not None:
|
||||
runs = [r for r in runs if r.kind == kind]
|
||||
if status is not None:
|
||||
runs = [r for r in runs if r.status == status]
|
||||
runs.sort(key=lambda r: r.created_at, reverse=True)
|
||||
total = len(runs)
|
||||
return runs[offset : offset + limit], total
|
||||
|
||||
|
||||
def get_run(run_id: str) -> BenchmarkRun | None:
|
||||
return _runs.get(run_id)
|
||||
|
||||
|
||||
def get_report(run_id: str) -> BenchmarkReport | None:
|
||||
return _reports.get(run_id)
|
||||
|
||||
|
||||
def get_events(run_id: str) -> list[BenchmarkEvent]:
|
||||
return _events.get(run_id, [])
|
||||
|
||||
|
||||
def cancel_run(run_id: str) -> BenchmarkRun | None:
|
||||
"""取消运行:对 queued/running 设置取消标志,后台 Task 在 Case 边界检查后置为 cancelled。"""
|
||||
run = _runs.get(run_id)
|
||||
if run is None:
|
||||
return None
|
||||
if run.status in (BenchmarkStatus.queued, BenchmarkStatus.running):
|
||||
_cancel_flags[run_id].set()
|
||||
return run
|
||||
|
||||
|
||||
def subscribe(run_id: str) -> asyncio.Queue[BenchmarkEvent] | None:
|
||||
"""订阅运行事件流;运行已结束(completed/failed/cancelled)时返回 None。"""
|
||||
run = _runs.get(run_id)
|
||||
if run is None or run.status in (
|
||||
BenchmarkStatus.completed,
|
||||
BenchmarkStatus.failed,
|
||||
BenchmarkStatus.cancelled,
|
||||
):
|
||||
return None
|
||||
queue: asyncio.Queue[BenchmarkEvent] = asyncio.Queue()
|
||||
_subscribers.setdefault(run_id, []).append(queue)
|
||||
return queue
|
||||
|
||||
|
||||
def unsubscribe(run_id: str, queue: asyncio.Queue[BenchmarkEvent]) -> None:
|
||||
subscribers = _subscribers.get(run_id)
|
||||
if subscribers and queue in subscribers:
|
||||
subscribers.remove(queue)
|
||||
|
||||
|
||||
async def wait_for_run(run_id: str) -> BenchmarkRun:
|
||||
"""等待后台任务结束(测试/轮询用);无任务时直接返回当前状态。"""
|
||||
task = _tasks.get(run_id)
|
||||
if task is not None:
|
||||
await task
|
||||
return _runs.get(run_id)
|
||||
@@ -24,6 +24,7 @@ class Settings:
|
||||
db_path: Path
|
||||
vault_path: Path
|
||||
attachments_path: Path
|
||||
benchmark_datasets_path: Path
|
||||
|
||||
|
||||
@lru_cache
|
||||
@@ -41,4 +42,7 @@ def get_settings() -> Settings:
|
||||
attachments_path=Path(
|
||||
os.getenv("APP_ATTACHMENTS_PATH", str(data_dir / "attachments"))
|
||||
),
|
||||
benchmark_datasets_path=Path(
|
||||
os.getenv("APP_BENCHMARK_DATASETS_PATH", str(data_dir / "benchmarks"))
|
||||
),
|
||||
)
|
||||
|
||||
@@ -5,7 +5,9 @@ from app.agent.builtin_tools import register_builtin_tools
|
||||
from app.contracts import ModelCapability, ProviderConfig, ProviderType
|
||||
from app.config import BACKEND_DIR, get_settings
|
||||
from app.extensions import PluginRuntime, SkillRuntime
|
||||
from app.extensions.mcp_registry import McpServerRegistry
|
||||
from app.providers import MockProvider, ProviderFactory, ProviderRegistry
|
||||
from app.providers.routing import ModelRoutingService
|
||||
from app.providers.credentials import (
|
||||
ChainedCredentialResolver,
|
||||
EncryptedCredentialStore,
|
||||
@@ -17,11 +19,13 @@ from app.providers.credentials import (
|
||||
class ApplicationContainer:
|
||||
providers: ProviderRegistry
|
||||
provider_factory: ProviderFactory
|
||||
model_routing: ModelRoutingService
|
||||
credentials: EncryptedCredentialStore
|
||||
tools: ToolRegistry
|
||||
permissions: PermissionManager
|
||||
skills: SkillRuntime
|
||||
plugins: PluginRuntime
|
||||
mcp_servers: McpServerRegistry
|
||||
agent: AgentRuntime
|
||||
|
||||
|
||||
@@ -31,7 +35,7 @@ def build_container() -> ApplicationContainer:
|
||||
provider_factory = ProviderFactory(
|
||||
ChainedCredentialResolver(credentials, EnvironmentCredentialResolver())
|
||||
)
|
||||
providers = ProviderRegistry()
|
||||
providers = ProviderRegistry(provider_factory)
|
||||
providers.register(
|
||||
ProviderConfig(
|
||||
provider_id="mock",
|
||||
@@ -61,6 +65,14 @@ def build_container() -> ApplicationContainer:
|
||||
plugins.install(BACKEND_DIR / "extensions" / "plugins" / "text-tools")
|
||||
plugins.enable("text-tools")
|
||||
|
||||
mcp_servers = McpServerRegistry(
|
||||
tools,
|
||||
credentials,
|
||||
settings.data_dir,
|
||||
allow_process_launch=settings.environment == "development",
|
||||
)
|
||||
mcp_servers.restore_enabled()
|
||||
|
||||
skills = SkillRuntime(tools)
|
||||
skills.install(BACKEND_DIR / "extensions" / "skills" / "knowledge-assistant")
|
||||
skills.enable("knowledge-assistant")
|
||||
@@ -76,11 +88,13 @@ def build_container() -> ApplicationContainer:
|
||||
return ApplicationContainer(
|
||||
providers=providers,
|
||||
provider_factory=provider_factory,
|
||||
model_routing=ModelRoutingService(providers, provider_factory.credentials),
|
||||
credentials=credentials,
|
||||
tools=tools,
|
||||
permissions=permissions,
|
||||
skills=skills,
|
||||
plugins=plugins,
|
||||
mcp_servers=mcp_servers,
|
||||
agent=agent,
|
||||
)
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ from datetime import datetime
|
||||
from enum import Enum
|
||||
from typing import Annotated, Any, Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, SecretStr
|
||||
from pydantic import BaseModel, ConfigDict, Field, SecretStr, field_validator
|
||||
|
||||
|
||||
class Contract(BaseModel):
|
||||
@@ -144,6 +144,12 @@ class SearchRequest(Contract):
|
||||
limit: int = Field(default=20, ge=1, le=100)
|
||||
offset: int = Field(default=0, ge=0)
|
||||
include_snippet: bool = True
|
||||
# 检索调优参数(Benchmark 与 Skill 共用):控制 RRF / 精排 / 候选池 / 分数阈值。
|
||||
# rerank_candidates=None 表示对全部候选精排(保留原有行为),Benchmark 传显式值。
|
||||
rrf_k: int = Field(default=60, ge=1)
|
||||
rerank: bool = True
|
||||
rerank_candidates: int | None = Field(default=None, ge=1)
|
||||
score_threshold: float = Field(default=0.0, ge=0.0)
|
||||
|
||||
|
||||
class Citation(Contract):
|
||||
@@ -199,7 +205,7 @@ class ToolDefinition(Contract):
|
||||
description: str
|
||||
parameters: dict[str, Any] = Field(default_factory=dict)
|
||||
permission: str | None = None
|
||||
source: Literal["builtin", "plugin"] = "builtin"
|
||||
source: Literal["builtin", "plugin", "mcp_server"] = "builtin"
|
||||
|
||||
|
||||
class ToolCall(Contract):
|
||||
@@ -230,6 +236,8 @@ class ModelCapability(str, Enum):
|
||||
streaming = "streaming"
|
||||
structured_output = "structured_output"
|
||||
embedding = "embedding"
|
||||
transcription = "transcription"
|
||||
speaker_matching = "speaker_matching"
|
||||
|
||||
|
||||
class ModelRequest(Contract):
|
||||
@@ -486,6 +494,97 @@ class PluginHostStatus(Contract):
|
||||
error: str | None = None
|
||||
|
||||
|
||||
# Independent user-managed MCP Server Registry. This is deliberately separate
|
||||
# from Plugin manifests: a server can contribute tools without being a Plugin.
|
||||
class McpServerTransport(str, Enum):
|
||||
stdio = "stdio"
|
||||
streamable_http = "streamable_http"
|
||||
sse = "sse"
|
||||
|
||||
|
||||
class McpServerConfig(Contract):
|
||||
name: str = Field(min_length=1, max_length=80)
|
||||
transport: McpServerTransport = McpServerTransport.stdio
|
||||
command: str | None = Field(default=None, max_length=1024)
|
||||
args: list[str] = Field(default_factory=list, max_length=64)
|
||||
url: str | None = Field(default=None, max_length=4096)
|
||||
headers: dict[str, str] = Field(default_factory=dict)
|
||||
environment: dict[str, str] = Field(default_factory=dict)
|
||||
secret_environment_keys: list[str] = Field(default_factory=list)
|
||||
secret_header_keys: list[str] = Field(default_factory=list)
|
||||
permissions: list[str] = Field(default_factory=list)
|
||||
startup_timeout_seconds: float = Field(default=15, ge=1, le=120)
|
||||
tool_timeout_seconds: float = Field(default=30, ge=1, le=300)
|
||||
|
||||
|
||||
class McpServerCreateRequest(McpServerConfig):
|
||||
pass
|
||||
|
||||
|
||||
class McpServerUpdateRequest(McpServerConfig):
|
||||
version: int = Field(ge=1)
|
||||
|
||||
|
||||
class McpServerSecretWriteRequest(Contract):
|
||||
secret: SecretStr = Field(min_length=1, max_length=32768)
|
||||
|
||||
|
||||
class McpServerSecretStatus(Contract):
|
||||
key: str
|
||||
configured: bool
|
||||
|
||||
|
||||
class McpServerTrustRequest(Contract):
|
||||
command_digest: str = Field(min_length=64, max_length=64)
|
||||
|
||||
|
||||
class McpServerStatus(Contract):
|
||||
enabled: bool = False
|
||||
status: PluginHostState = PluginHostState.stopped
|
||||
tools_count: int = 0
|
||||
protocol_version: str | None = None
|
||||
remote_server_name: str | None = None
|
||||
remote_server_version: str | None = None
|
||||
error: str | None = None
|
||||
last_tested_at: datetime | None = None
|
||||
last_test_succeeded: bool | None = None
|
||||
|
||||
|
||||
class McpServer(McpServerStatus):
|
||||
server_id: str
|
||||
version: int
|
||||
name: str
|
||||
transport: McpServerTransport
|
||||
command: str | None = None
|
||||
args: list[str] = Field(default_factory=list)
|
||||
url: str | None = None
|
||||
headers: dict[str, str] = Field(default_factory=dict)
|
||||
environment: dict[str, str] = Field(default_factory=dict)
|
||||
permissions: list[str] = Field(default_factory=list)
|
||||
startup_timeout_seconds: float
|
||||
tool_timeout_seconds: float
|
||||
secret_environment: dict[str, bool] = Field(default_factory=dict)
|
||||
secret_headers: dict[str, bool] = Field(default_factory=dict)
|
||||
trusted: bool = False
|
||||
command_digest: str
|
||||
command_summary: str
|
||||
|
||||
|
||||
class McpServerListResponse(Contract):
|
||||
items: list[McpServer] = Field(default_factory=list)
|
||||
|
||||
|
||||
class McpToolSummary(Contract):
|
||||
name: str
|
||||
remote_name: str
|
||||
description: str
|
||||
permission: str | None = None
|
||||
|
||||
|
||||
class McpToolSummaryListResponse(Contract):
|
||||
items: list[McpToolSummary] = Field(default_factory=list)
|
||||
|
||||
|
||||
class PluginCommandLocation(str, Enum):
|
||||
command_palette = "command_palette"
|
||||
context_menu = "context_menu"
|
||||
@@ -666,7 +765,24 @@ class ProviderType(str, Enum):
|
||||
ollama = "ollama"
|
||||
|
||||
|
||||
class ProviderConfig(Contract):
|
||||
class ProviderConnectionFields(Contract):
|
||||
base_url: str | None = None
|
||||
credential_id: str | None = None
|
||||
|
||||
@field_validator("base_url")
|
||||
@classmethod
|
||||
def provider_url(cls, value: str | None) -> str | None:
|
||||
if value is None:
|
||||
return value
|
||||
from urllib.parse import urlsplit
|
||||
parsed = urlsplit(value)
|
||||
if (parsed.scheme not in {"http", "https"} or not parsed.hostname or
|
||||
parsed.username or parsed.password or parsed.query or parsed.fragment):
|
||||
raise ValueError("Base URL requires HTTP(S), without credentials, query or fragment")
|
||||
return value.rstrip("/")
|
||||
|
||||
|
||||
class ProviderConfig(ProviderConnectionFields):
|
||||
provider_id: str
|
||||
provider_type: ProviderType
|
||||
name: str
|
||||
@@ -677,7 +793,7 @@ class ProviderConfig(Contract):
|
||||
capabilities: list[ModelCapability] = Field(default_factory=list)
|
||||
|
||||
|
||||
class ProviderCreateRequest(Contract):
|
||||
class ProviderCreateRequest(ProviderConnectionFields):
|
||||
provider_type: ProviderType
|
||||
name: str
|
||||
base_url: str | None = None
|
||||
@@ -686,7 +802,8 @@ class ProviderCreateRequest(Contract):
|
||||
enabled: bool = True
|
||||
|
||||
|
||||
class ProviderUpdateRequest(Contract):
|
||||
class ProviderUpdateRequest(ProviderConnectionFields):
|
||||
provider_type: ProviderType | None = None
|
||||
name: str | None = None
|
||||
base_url: str | None = None
|
||||
default_model: str | None = None
|
||||
@@ -705,6 +822,80 @@ class ProviderPreset(Contract):
|
||||
base_url: str
|
||||
default_credential_id: str | None = None
|
||||
requires_credential: bool = True
|
||||
logo_id: str = "custom"
|
||||
description: str = ""
|
||||
capabilities: list[ModelCapability] = Field(default_factory=list)
|
||||
|
||||
|
||||
class ModelBinding(Contract):
|
||||
provider_id: str = Field(min_length=1, max_length=128)
|
||||
model: str = Field(min_length=1, max_length=256)
|
||||
endpoint: str = Field(min_length=1, max_length=256)
|
||||
dimensions: int | None = Field(default=None, ge=1, le=16384)
|
||||
|
||||
@field_validator("endpoint")
|
||||
@classmethod
|
||||
def relative_endpoint(cls, value: str) -> str:
|
||||
# An endpoint is a path on the selected provider, never a second origin.
|
||||
import re
|
||||
if not re.fullmatch(r"/[A-Za-z0-9_/-]+", value) or value.startswith("//"):
|
||||
raise ValueError("endpoint must be an absolute API path on the provider")
|
||||
return value
|
||||
|
||||
@field_validator("model", "provider_id")
|
||||
@classmethod
|
||||
def non_blank(cls, value: str) -> str:
|
||||
if not value.strip():
|
||||
raise ValueError("value must not be blank")
|
||||
return value.strip()
|
||||
|
||||
|
||||
class ModelRoutingConfig(Contract):
|
||||
version: int = Field(default=0, ge=0)
|
||||
embedding: ModelBinding | None = None
|
||||
transcription: ModelBinding | None = None
|
||||
speaker_matching: ModelBinding | None = None
|
||||
|
||||
|
||||
class LocalBackendStatus(Contract):
|
||||
capability: Literal["embedding", "transcription", "speaker_matching"]
|
||||
status: Literal["placeholder", "not_installed", "ready"]
|
||||
message: str
|
||||
|
||||
|
||||
class ModelRoutingResponse(Contract):
|
||||
config: ModelRoutingConfig
|
||||
local_backends: list[LocalBackendStatus]
|
||||
|
||||
|
||||
class EmbeddingRequest(Contract):
|
||||
texts: list[str] = Field(min_length=1, max_length=256)
|
||||
|
||||
@field_validator("texts")
|
||||
@classmethod
|
||||
def bound_texts(cls, value: list[str]) -> list[str]:
|
||||
if sum(len(text) for text in value) > 200_000:
|
||||
raise ValueError("embedding input is too large")
|
||||
return value
|
||||
|
||||
|
||||
class EmbeddingResult(Contract):
|
||||
vectors: list[list[float]]
|
||||
source: Literal["api", "local"]
|
||||
model_id: str
|
||||
dimensions: int
|
||||
fallback_reason: str | None = None
|
||||
|
||||
|
||||
class SpeakerMatchRequest(Contract):
|
||||
attachment_id: str
|
||||
reference_attachment_id: str
|
||||
|
||||
|
||||
class SpeakerMatchResult(Contract):
|
||||
score: float = Field(ge=0, le=1, allow_inf_nan=False)
|
||||
source: Literal["api", "local"]
|
||||
fallback_reason: str | None = None
|
||||
|
||||
|
||||
class ProviderPresetListResponse(Contract):
|
||||
@@ -797,6 +988,8 @@ class TranscriptionJob(Contract):
|
||||
error_code: str | None = None
|
||||
error_message: str | None = None
|
||||
created_at: datetime
|
||||
source: Literal["api", "local", "sidecar"] | None = None
|
||||
fallback_reason: str | None = None
|
||||
|
||||
|
||||
class IndexStatus(Contract):
|
||||
@@ -818,3 +1011,152 @@ class IndexJob(Contract):
|
||||
status: Literal["queued", "running", "completed", "failed"]
|
||||
scope: Literal["all", "notes", "vectors"]
|
||||
created_at: datetime
|
||||
|
||||
|
||||
# Benchmark
|
||||
class BenchmarkKind(str, Enum):
|
||||
rag = "rag"
|
||||
agent = "agent"
|
||||
|
||||
|
||||
class BenchmarkStatus(str, Enum):
|
||||
queued = "queued"
|
||||
running = "running"
|
||||
completed = "completed"
|
||||
failed = "failed"
|
||||
cancelled = "cancelled"
|
||||
|
||||
|
||||
class RAGDatasetCase(Contract):
|
||||
case_id: str
|
||||
query: str = Field(min_length=1)
|
||||
expected_note_ids: list[str] = Field(default_factory=list)
|
||||
expected_block_ids: list[str] = Field(default_factory=list)
|
||||
citation_required: bool = False
|
||||
tags: list[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class RAGRetrievalConfig(Contract):
|
||||
"""RAG Benchmark 的检索参数。top_k 映射到 SearchRequest.limit,
|
||||
其余参数透传到 SearchRequest,由检索引擎实际执行。"""
|
||||
|
||||
top_k: int = Field(default=10, ge=1, le=100)
|
||||
rrf_k: int = Field(default=60, ge=1)
|
||||
rerank: bool = True
|
||||
rerank_candidates: int = Field(default=20, ge=1)
|
||||
score_threshold: float = Field(default=0.0, ge=0.0)
|
||||
|
||||
|
||||
class RAGRunRequest(Contract):
|
||||
dataset_id: str = Field(min_length=1)
|
||||
modes: list[SearchMode] = Field(
|
||||
default_factory=lambda: [SearchMode.fts, SearchMode.vector, SearchMode.hybrid],
|
||||
min_length=1,
|
||||
)
|
||||
retrieval: RAGRetrievalConfig = Field(default_factory=RAGRetrievalConfig)
|
||||
repeat: int = Field(default=1, ge=1, le=10)
|
||||
metadata: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
@field_validator("modes")
|
||||
@classmethod
|
||||
def _no_duplicate_modes(cls, value: list[SearchMode]) -> list[SearchMode]:
|
||||
if len(value) != len(set(value)):
|
||||
raise ValueError("modes must not contain duplicates")
|
||||
return value
|
||||
|
||||
|
||||
class RAGMetrics(Contract):
|
||||
hit_at_1: float = 0.0
|
||||
hit_at_5: float = 0.0
|
||||
recall_at_k: float = 0.0
|
||||
mrr: float = 0.0
|
||||
citation_hit_rate: float = 0.0
|
||||
p50_latency_ms: float = 0.0
|
||||
p95_latency_ms: float = 0.0
|
||||
# 样本构成:失败样本按零分计入质量指标,汇总不虚高;报告据此可知实际分母
|
||||
total_cases: int = 0
|
||||
successful_cases: int = 0
|
||||
failed_cases: int = 0
|
||||
failure_rate: float = 0.0
|
||||
|
||||
|
||||
class BenchmarkDatasetInfo(Contract):
|
||||
dataset_id: str
|
||||
kind: BenchmarkKind
|
||||
version: str
|
||||
description: str = ""
|
||||
case_count: int
|
||||
content_hash: str
|
||||
|
||||
|
||||
class BenchmarkDatasetListResponse(Contract):
|
||||
items: list[BenchmarkDatasetInfo] = Field(default_factory=list)
|
||||
|
||||
|
||||
class BenchmarkRun(Contract):
|
||||
run_id: str
|
||||
kind: BenchmarkKind
|
||||
dataset_id: str
|
||||
dataset_hash: str
|
||||
status: BenchmarkStatus
|
||||
progress: float | None = None
|
||||
metrics: dict[str, Any] | None = None
|
||||
config_snapshot: dict[str, Any] = Field(default_factory=dict)
|
||||
error: str | None = None
|
||||
error_code: str | None = None
|
||||
created_at: datetime
|
||||
started_at: datetime | None = None
|
||||
completed_at: datetime | None = None
|
||||
|
||||
|
||||
class BenchmarkRunListResponse(Contract):
|
||||
items: list[BenchmarkRun] = Field(default_factory=list)
|
||||
page: PageMeta = Field(default_factory=PageMeta)
|
||||
|
||||
|
||||
class BenchmarkEventType(str, Enum):
|
||||
run_started = "RunStarted"
|
||||
case_completed = "CaseCompleted"
|
||||
run_completed = "RunCompleted"
|
||||
run_failed = "RunFailed"
|
||||
run_cancelled = "RunCancelled"
|
||||
|
||||
|
||||
class BenchmarkEvent(Contract):
|
||||
event: BenchmarkEventType
|
||||
run_id: str
|
||||
sequence: int
|
||||
data: dict[str, Any] = Field(default_factory=dict)
|
||||
timestamp: datetime
|
||||
|
||||
|
||||
class RAGCaseResult(Contract):
|
||||
embedding: dict[str, Any] = Field(default_factory=dict)
|
||||
case_id: str
|
||||
mode: SearchMode
|
||||
repeat: int
|
||||
latency_ms: float
|
||||
retrieved_note_ids: list[str] = Field(default_factory=list)
|
||||
retrieved_block_ids: list[str] = Field(default_factory=list)
|
||||
hit_at_1: bool = False
|
||||
hit_at_5: bool = False
|
||||
recall: float = 0.0
|
||||
reciprocal_rank: float = 0.0
|
||||
citation_hit: bool = False
|
||||
# 该 Case 是否声明了 expected_block_ids(决定是否计入 citation_hit_rate 分母)
|
||||
citation_applicable: bool = False
|
||||
error: str | None = None
|
||||
error_code: str | None = None
|
||||
|
||||
|
||||
class BenchmarkReport(Contract):
|
||||
run_id: str
|
||||
kind: BenchmarkKind
|
||||
dataset_id: str
|
||||
dataset_hash: str
|
||||
status: BenchmarkStatus
|
||||
config_snapshot: dict[str, Any] = Field(default_factory=dict)
|
||||
metrics: dict[str, Any] = Field(default_factory=dict)
|
||||
cases: list[RAGCaseResult] = Field(default_factory=list)
|
||||
error: str | None = None
|
||||
error_code: str | None = None
|
||||
|
||||
@@ -36,7 +36,11 @@ async def validation_error_handler(_: Request, exc: RequestValidationError) -> J
|
||||
error=ErrorDetail(
|
||||
code="VALIDATION_ERROR",
|
||||
message="Request validation failed.",
|
||||
details={"errors": exc.errors()},
|
||||
# Pydantic ctx can contain exception objects; input may contain API keys.
|
||||
details={"errors": [
|
||||
{key: error[key] for key in ("type", "loc", "msg") if key in error}
|
||||
for error in exc.errors()
|
||||
]},
|
||||
)
|
||||
)
|
||||
return JSONResponse(status_code=422, content=jsonable_encoder(body))
|
||||
|
||||
@@ -19,6 +19,7 @@ async def lifespan(_: FastAPI):
|
||||
yield
|
||||
# 第三方 MCP Server 必须跟随 AI Core 退出,不能遗留孤儿进程。
|
||||
container.plugins.shutdown()
|
||||
container.mcp_servers.shutdown()
|
||||
|
||||
|
||||
app = FastAPI(
|
||||
|
||||
@@ -0,0 +1,163 @@
|
||||
"""Native Anthropic Messages protocol with incrementally decoded content blocks."""
|
||||
|
||||
import json
|
||||
from contextlib import aclosing
|
||||
|
||||
from app.contracts import MessageRole, ModelEventType, ModelRequest
|
||||
from app.providers.base import ProviderError, ProviderToolCall, ProviderTurn
|
||||
from app.providers.http_base import (
|
||||
UsageTracker, check_error, decode_tool_arguments, invalid_response, list_value,
|
||||
object_value, string_value, token_count, truncated_stream,
|
||||
)
|
||||
from app.providers.openai_compatible import OpenAICompatibleProvider
|
||||
from app.providers.tool_names import mapped_tool_names
|
||||
|
||||
|
||||
class AnthropicMessagesProvider(OpenAICompatibleProvider):
|
||||
stream_path = "/messages"
|
||||
|
||||
def _headers(self) -> dict[str, str]:
|
||||
headers = super()._headers()
|
||||
authorization = headers.pop("Authorization", None)
|
||||
if authorization:
|
||||
headers["x-api-key"] = authorization.removeprefix("Bearer ")
|
||||
headers["anthropic-version"] = "2023-06-01"
|
||||
return headers
|
||||
|
||||
def _payload(self, request: ModelRequest, *, stream: bool) -> dict[str, object]:
|
||||
systems = [request.system] if request.system else []
|
||||
messages = []
|
||||
for message in request.messages:
|
||||
if message.role == MessageRole.system:
|
||||
systems.append(message.content)
|
||||
continue
|
||||
if message.role == MessageRole.tool:
|
||||
if not message.tool_call_id:
|
||||
raise ProviderError("PROVIDER_INVALID_REQUEST", "Tool result requires a call identifier.")
|
||||
role = "user"
|
||||
content = [{"type": "tool_result", "tool_use_id": message.tool_call_id, "content": message.content}]
|
||||
else:
|
||||
role = message.role.value
|
||||
content = [{"type": "text", "text": message.content}] if message.content else []
|
||||
content += [{"type": "tool_use", "id": call.tool_call_id, "name": call.name,
|
||||
"input": call.arguments} for call in message.tool_calls]
|
||||
if not content:
|
||||
continue
|
||||
if messages and messages[-1]["role"] == role:
|
||||
messages[-1]["content"].extend(content)
|
||||
else:
|
||||
messages.append({"role": role, "content": content})
|
||||
payload: dict[str, object] = {"model": request.model, "messages": messages,
|
||||
"max_tokens": request.max_tokens or 4096, "stream": stream}
|
||||
if systems:
|
||||
payload["system"] = "\n\n".join(systems)
|
||||
if request.tools:
|
||||
payload["tools"] = [{"name": tool.name, "description": tool.description,
|
||||
"input_schema": tool.parameters} for tool in request.tools]
|
||||
if request.temperature is not None:
|
||||
payload["temperature"] = request.temperature
|
||||
if request.response_format is not None:
|
||||
format_ = request.response_format
|
||||
if format_.get("type") != "json_schema":
|
||||
raise ProviderError("PROVIDER_INVALID_REQUEST", "Messages requires a JSON schema response format.")
|
||||
schema = object_value(format_.get("json_schema"))
|
||||
payload["output_config"] = {"format": {"type": "json_schema", "schema": object_value(schema.get("schema"))}}
|
||||
return payload
|
||||
|
||||
@mapped_tool_names
|
||||
async def complete(self, request: ModelRequest) -> ProviderTurn:
|
||||
data = await self._request("POST", self.stream_path, json=self._payload(request, stream=False))
|
||||
texts = []
|
||||
calls = []
|
||||
for raw in list_value(data.get("content")):
|
||||
block = object_value(raw)
|
||||
if block.get("type") == "text":
|
||||
texts.append(string_value(block.get("text")))
|
||||
elif block.get("type") == "tool_use":
|
||||
calls.append(ProviderToolCall(
|
||||
tool_call_id=string_value(block.get("id"), nonempty=True),
|
||||
name=string_value(block.get("name"), nonempty=True),
|
||||
arguments=decode_tool_arguments(block.get("input")),
|
||||
))
|
||||
return ProviderTurn(text="".join(texts) or None, tool_calls=calls,
|
||||
**UsageTracker(cache_tokens=True).update(data.get("usage") or {}))
|
||||
|
||||
async def _events(self, request: ModelRequest):
|
||||
blocks: dict[int, dict] = {}
|
||||
usage = UsageTracker(cache_tokens=True)
|
||||
started = False
|
||||
async with aclosing(self._stream_json(self._payload(request, stream=True))) as chunks:
|
||||
async for data in chunks:
|
||||
kind = string_value(data.get("type"), nonempty=True)
|
||||
if kind == "message_start":
|
||||
if started:
|
||||
raise invalid_response()
|
||||
started = True
|
||||
message = object_value(data.get("message"))
|
||||
check_error(message)
|
||||
if message.get("usage") is not None:
|
||||
yield ModelEventType.usage, usage.update(message["usage"])
|
||||
elif kind == "content_block_start":
|
||||
index = token_count(data.get("index"))
|
||||
if not started or index in blocks:
|
||||
raise invalid_response()
|
||||
block = dict(object_value(data.get("content_block")))
|
||||
blocks[index] = block
|
||||
block["closed"] = False
|
||||
if block.get("type") == "tool_use":
|
||||
block["id"] = string_value(block.get("id"), nonempty=True)
|
||||
block["name"] = string_value(block.get("name"), nonempty=True)
|
||||
block["arguments"] = ""
|
||||
block["input"] = object_value(block.get("input", {}))
|
||||
yield ModelEventType.tool_call_start, {"tool_call_id": block["id"], "name": block["name"]}
|
||||
elif block.get("type") == "text" and block.get("text"):
|
||||
yield ModelEventType.text_delta, {"text": string_value(block["text"])}
|
||||
elif block.get("type") == "thinking" and block.get("thinking"):
|
||||
yield ModelEventType.thinking_delta, {"text": string_value(block["thinking"])}
|
||||
elif kind == "content_block_delta":
|
||||
block = blocks.get(token_count(data.get("index")))
|
||||
if block is None or block["closed"]:
|
||||
raise invalid_response()
|
||||
delta = object_value(data.get("delta"))
|
||||
delta_type = delta.get("type")
|
||||
if delta_type == "text_delta":
|
||||
if block.get("type") != "text":
|
||||
raise invalid_response()
|
||||
yield ModelEventType.text_delta, {"text": string_value(delta.get("text"))}
|
||||
elif delta_type == "thinking_delta":
|
||||
if block.get("type") != "thinking":
|
||||
raise invalid_response()
|
||||
yield ModelEventType.thinking_delta, {"text": string_value(delta.get("thinking"))}
|
||||
elif delta_type == "input_json_delta" and block.get("type") == "tool_use":
|
||||
fragment = string_value(delta.get("partial_json"))
|
||||
block["arguments"] += fragment
|
||||
yield ModelEventType.tool_call_delta, {"tool_call_id": block["id"], "arguments_delta": fragment}
|
||||
# Signatures and future delta types have no representation in ModelEvent.
|
||||
elif kind == "content_block_stop":
|
||||
block = blocks.get(token_count(data.get("index")))
|
||||
if block is None or block["closed"]:
|
||||
raise invalid_response()
|
||||
block["closed"] = True
|
||||
if block.get("type") == "tool_use":
|
||||
if block["arguments"]:
|
||||
decode_tool_arguments(block["arguments"])
|
||||
else:
|
||||
yield ModelEventType.tool_call_delta, {
|
||||
"tool_call_id": block["id"], "arguments_delta": json.dumps(block["input"]),
|
||||
}
|
||||
yield ModelEventType.tool_call_end, {"tool_call_id": block["id"]}
|
||||
elif kind == "message_delta":
|
||||
if not started:
|
||||
raise invalid_response()
|
||||
object_value(data.get("delta"))
|
||||
if data.get("usage") is not None:
|
||||
yield ModelEventType.usage, usage.update(data["usage"])
|
||||
elif kind == "message_stop":
|
||||
if not started:
|
||||
raise invalid_response()
|
||||
if any(not block["closed"] for block in blocks.values()):
|
||||
raise truncated_stream()
|
||||
return
|
||||
elif kind == "[DONE]":
|
||||
raise truncated_stream()
|
||||
raise truncated_stream()
|
||||
@@ -5,15 +5,15 @@ import os
|
||||
import re
|
||||
import threading
|
||||
from pathlib import Path
|
||||
from typing import Protocol
|
||||
from typing import ClassVar, Protocol
|
||||
|
||||
from cryptography.fernet import Fernet, InvalidToken
|
||||
|
||||
from app.config import get_settings
|
||||
|
||||
|
||||
_CREDENTIAL_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]{0,127}$")
|
||||
_PLUGIN_CREDENTIAL_PREFIX = "plugin."
|
||||
_MCP_CREDENTIAL_PREFIX = "mcp."
|
||||
|
||||
|
||||
class CredentialStoreError(RuntimeError):
|
||||
@@ -27,16 +27,18 @@ class CredentialResolver(Protocol):
|
||||
def validate_provider_credential_id(credential_id: str | None) -> None:
|
||||
"""阻止 Provider 和通用凭据 API 跨入 Plugin 私有命名空间。"""
|
||||
|
||||
if credential_id and credential_id.casefold().startswith(
|
||||
_PLUGIN_CREDENTIAL_PREFIX
|
||||
):
|
||||
raise CredentialStoreError("Credential namespace is reserved for Plugin settings.")
|
||||
if credential_id and credential_id.casefold().startswith(_PLUGIN_CREDENTIAL_PREFIX):
|
||||
raise CredentialStoreError(
|
||||
"Credential namespace is reserved for Plugin settings."
|
||||
)
|
||||
if credential_id and credential_id.casefold().startswith(_MCP_CREDENTIAL_PREFIX):
|
||||
raise CredentialStoreError("Credential namespace is reserved for MCP settings.")
|
||||
|
||||
|
||||
class EnvironmentCredentialResolver:
|
||||
"""解析由桌面 Host 注入 Sidecar 进程的临时凭证上下文。"""
|
||||
|
||||
_development_aliases = {
|
||||
_development_aliases: ClassVar[dict[str, str]] = {
|
||||
"openai": "OPENAI_API_KEY",
|
||||
"deepseek": "DEEPSEEK_API_KEY",
|
||||
}
|
||||
@@ -84,7 +86,9 @@ class EncryptedCredentialStore:
|
||||
try:
|
||||
return Fernet(environment_key.encode("ascii"))
|
||||
except (ValueError, UnicodeEncodeError) as exc:
|
||||
raise CredentialStoreError("APP_CREDENTIAL_MASTER_KEY is invalid.") from exc
|
||||
raise CredentialStoreError(
|
||||
"APP_CREDENTIAL_MASTER_KEY is invalid."
|
||||
) from exc
|
||||
|
||||
key_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
self._restrict(key_path.parent, 0o700)
|
||||
@@ -101,7 +105,9 @@ class EncryptedCredentialStore:
|
||||
try:
|
||||
return Fernet(key_path.read_bytes().strip())
|
||||
except (OSError, ValueError) as exc:
|
||||
raise CredentialStoreError("Credential master key cannot be loaded.") from exc
|
||||
raise CredentialStoreError(
|
||||
"Credential master key cannot be loaded."
|
||||
) from exc
|
||||
|
||||
def _read_tokens(self) -> dict[str, str]:
|
||||
_, store_path = self._paths()
|
||||
@@ -110,11 +116,16 @@ class EncryptedCredentialStore:
|
||||
try:
|
||||
data = json.loads(store_path.read_text(encoding="utf-8"))
|
||||
except (OSError, json.JSONDecodeError) as exc:
|
||||
raise CredentialStoreError("Encrypted credential store cannot be loaded.") from exc
|
||||
raise CredentialStoreError(
|
||||
"Encrypted credential store cannot be loaded."
|
||||
) from exc
|
||||
if not isinstance(data, dict) or not all(
|
||||
isinstance(key, str) and isinstance(value, str) for key, value in data.items()
|
||||
isinstance(key, str) and isinstance(value, str)
|
||||
for key, value in data.items()
|
||||
):
|
||||
raise CredentialStoreError("Encrypted credential store has an invalid format.")
|
||||
raise CredentialStoreError(
|
||||
"Encrypted credential store has an invalid format."
|
||||
)
|
||||
return data
|
||||
|
||||
def _write_tokens(self, tokens: dict[str, str]) -> None:
|
||||
@@ -195,6 +206,22 @@ class EncryptedCredentialStore:
|
||||
self._write_tokens(tokens)
|
||||
return removed
|
||||
|
||||
def move_many(self, replacements: dict[str, str]) -> None:
|
||||
"""原子迁移凭据 ID,直接移动密文且不覆盖已经写入的新凭据。"""
|
||||
|
||||
for old_id, new_id in replacements.items():
|
||||
self._validate_id(old_id)
|
||||
self._validate_id(new_id)
|
||||
with self._lock:
|
||||
tokens = self._read_tokens()
|
||||
changed = False
|
||||
for old_id, new_id in replacements.items():
|
||||
if old_id != new_id and old_id in tokens:
|
||||
tokens.setdefault(new_id, tokens.pop(old_id))
|
||||
changed = True
|
||||
if changed:
|
||||
self._write_tokens(tokens)
|
||||
|
||||
|
||||
class ChainedCredentialResolver:
|
||||
def __init__(self, *resolvers: CredentialResolver) -> None:
|
||||
|
||||
@@ -16,6 +16,18 @@ class ProviderFactory:
|
||||
self.credentials = ProviderCredentialResolver(credentials)
|
||||
|
||||
def build(self, config: ProviderConfig) -> ModelProvider:
|
||||
if config.provider_type == ProviderType.openai_responses:
|
||||
from app.providers.openai_responses import OpenAIResponsesProvider
|
||||
return OpenAIResponsesProvider(
|
||||
base_url=config.base_url or "https://api.openai.com/v1",
|
||||
credential_id=config.credential_id, credentials=self.credentials,
|
||||
)
|
||||
if config.provider_type == ProviderType.anthropic_messages:
|
||||
from app.providers.anthropic_messages import AnthropicMessagesProvider
|
||||
return AnthropicMessagesProvider(
|
||||
base_url=config.base_url or "https://api.anthropic.com/v1",
|
||||
credential_id=config.credential_id, credentials=self.credentials,
|
||||
)
|
||||
if config.provider_type in {
|
||||
ProviderType.openai_chat,
|
||||
ProviderType.openai_compatible,
|
||||
@@ -31,7 +43,7 @@ class ProviderFactory:
|
||||
|
||||
@staticmethod
|
||||
def presets() -> list[ProviderPreset]:
|
||||
return [
|
||||
presets = [
|
||||
ProviderPreset(
|
||||
preset_id="openai",
|
||||
name="OpenAI",
|
||||
@@ -54,12 +66,45 @@ class ProviderFactory:
|
||||
requires_credential=False,
|
||||
),
|
||||
]
|
||||
# General API endpoints. Coding-plan endpoints and keys are separate products.
|
||||
domestic = [
|
||||
("kimi", "Kimi / 月之暗面", "https://api.moonshot.cn/v1", [], "长上下文对话;模型以账号权限为准。"),
|
||||
("qwen", "阿里云百炼", "https://dashscope.aliyuncs.com/compatible-mode/v1", [ModelCapability.embedding], "中国内地兼容接口;海外地域需修改地址。"),
|
||||
("zhipu", "智谱 GLM", "https://open.bigmodel.cn/api/paas/v4", [ModelCapability.embedding], "通用 API;Coding Plan 请使用其专用地址。"),
|
||||
("volcengine", "火山方舟 / 豆包", "https://ark.cn-beijing.volces.com/api/v3", [ModelCapability.embedding], "按账号填写模型 ID 或推理接入点 ID。"),
|
||||
("siliconflow", "硅基流动", "https://api.siliconflow.cn/v1", [ModelCapability.embedding, ModelCapability.transcription], "支持兼容 Embedding 和音频转写接口。"),
|
||||
("baidu", "百度千帆", "https://qianfan.baidubce.com/v2", [ModelCapability.embedding], "使用千帆 API Key;模型列表取决于账号。"),
|
||||
("hunyuan", "腾讯混元", "https://api.hunyuan.cloud.tencent.com/v1", [], "OpenAI 兼容对话接口。"),
|
||||
("minimax", "MiniMax", "https://api.minimaxi.com/v1", [], "文本对话兼容接口;其他媒体协议需独立适配。"),
|
||||
("stepfun", "阶跃星辰", "https://api.stepfun.com/v1", [], "通用 API;Step Plan 请使用其专用地址。"),
|
||||
]
|
||||
for preset_id, name, url, extra, description in domestic:
|
||||
presets.append(ProviderPreset(
|
||||
preset_id=preset_id, name=name, provider_type=ProviderType.openai_compatible,
|
||||
base_url=url, default_credential_id=preset_id, logo_id=preset_id,
|
||||
capabilities=[ModelCapability.chat, *extra], description=description,
|
||||
))
|
||||
presets.extend([
|
||||
ProviderPreset(preset_id="openai-responses", name="OpenAI Responses", provider_type=ProviderType.openai_responses,
|
||||
base_url="https://api.openai.com/v1", default_credential_id="openai", logo_id="openai"),
|
||||
ProviderPreset(preset_id="anthropic", name="Anthropic / Claude", provider_type=ProviderType.anthropic_messages,
|
||||
base_url="https://api.anthropic.com/v1", default_credential_id="anthropic", logo_id="anthropic"),
|
||||
])
|
||||
for preset in presets:
|
||||
if preset.logo_id == "custom":
|
||||
preset.logo_id = preset.preset_id
|
||||
if not preset.capabilities:
|
||||
preset.capabilities = [ModelCapability.chat]
|
||||
presets[0].capabilities += [ModelCapability.embedding, ModelCapability.transcription]
|
||||
return presets
|
||||
|
||||
@staticmethod
|
||||
def capabilities(provider_type: ProviderType) -> list[ModelCapability]:
|
||||
if provider_type in {
|
||||
ProviderType.openai_chat,
|
||||
ProviderType.openai_compatible,
|
||||
ProviderType.openai_responses,
|
||||
ProviderType.anthropic_messages,
|
||||
}:
|
||||
return [
|
||||
ModelCapability.chat,
|
||||
|
||||
@@ -1,9 +1,13 @@
|
||||
import json
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import aclosing
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import httpx
|
||||
|
||||
from app.contracts import ModelEvent, ModelEventType, ModelRequest
|
||||
from app.providers.base import ProviderError, ProviderTurn
|
||||
from app.providers.tool_names import prepare_tool_names
|
||||
|
||||
|
||||
class TurnStreamingMixin:
|
||||
@@ -80,3 +84,229 @@ def decode_tool_arguments(value: object) -> dict[str, object]:
|
||||
if not isinstance(decoded, dict):
|
||||
raise ProviderError("PROVIDER_INVALID_RESPONSE", "Tool arguments must be an object.")
|
||||
return decoded
|
||||
|
||||
|
||||
def invalid_response() -> ProviderError:
|
||||
return ProviderError("PROVIDER_INVALID_RESPONSE", "Provider returned an invalid response.")
|
||||
|
||||
|
||||
def truncated_stream() -> ProviderError:
|
||||
return ProviderError("PROVIDER_STREAM_TRUNCATED", "Provider stream ended before completion.")
|
||||
|
||||
|
||||
def object_value(value: object) -> dict:
|
||||
if not isinstance(value, dict):
|
||||
raise invalid_response()
|
||||
return value
|
||||
|
||||
|
||||
def list_value(value: object) -> list:
|
||||
if not isinstance(value, list):
|
||||
raise invalid_response()
|
||||
return value
|
||||
|
||||
|
||||
def string_value(value: object, *, nonempty: bool = False) -> str:
|
||||
if not isinstance(value, str) or (nonempty and not value):
|
||||
raise invalid_response()
|
||||
return value
|
||||
|
||||
|
||||
def token_count(value: object) -> int:
|
||||
if isinstance(value, bool) or not isinstance(value, int) or value < 0:
|
||||
raise invalid_response()
|
||||
return value
|
||||
|
||||
|
||||
def remote_error(value: object) -> ProviderError:
|
||||
# Never reflect upstream messages, URLs, request bodies or credentials.
|
||||
error = value if isinstance(value, dict) else {}
|
||||
code = error.get("code") or error.get("type")
|
||||
mapping = {
|
||||
"authentication_error": "PROVIDER_AUTH_FAILED",
|
||||
"invalid_api_key": "PROVIDER_AUTH_FAILED",
|
||||
"permission_error": "PROVIDER_AUTH_FAILED",
|
||||
"rate_limit_error": "PROVIDER_RATE_LIMITED",
|
||||
"rate_limit_exceeded": "PROVIDER_RATE_LIMITED",
|
||||
"insufficient_quota": "PROVIDER_RATE_LIMITED",
|
||||
"not_found_error": "MODEL_NOT_FOUND",
|
||||
"model_not_found": "MODEL_NOT_FOUND",
|
||||
"invalid_request_error": "PROVIDER_INVALID_REQUEST",
|
||||
"context_length_exceeded": "PROVIDER_INVALID_REQUEST",
|
||||
}
|
||||
mapped = mapping.get(code, "PROVIDER_UNAVAILABLE") if isinstance(code, str) else "PROVIDER_UNAVAILABLE"
|
||||
return ProviderError(mapped, "Provider could not complete the request.")
|
||||
|
||||
|
||||
def check_error(data: dict) -> None:
|
||||
if data.get("error") is not None or data.get("type") == "error":
|
||||
raise remote_error(data.get("error") or data)
|
||||
|
||||
|
||||
class UsageTracker:
|
||||
"""Merge cumulative snapshots, including partial usage updates."""
|
||||
|
||||
def __init__(self, input_key: str = "input_tokens", output_key: str = "output_tokens",
|
||||
*, cache_tokens: bool = False) -> None:
|
||||
self.input_key = input_key
|
||||
self.output_key = output_key
|
||||
self.cache_tokens = cache_tokens
|
||||
self.counts: dict[str, int] = {}
|
||||
|
||||
def update(self, value: object) -> dict[str, int]:
|
||||
usage = object_value(value)
|
||||
keys = [self.input_key, self.output_key]
|
||||
if self.cache_tokens:
|
||||
keys += ["cache_creation_input_tokens", "cache_read_input_tokens"]
|
||||
for key in keys:
|
||||
if key in usage:
|
||||
self.counts[key] = max(self.counts.get(key, 0), token_count(usage[key]))
|
||||
inputs = self.counts.get(self.input_key, 0)
|
||||
if self.cache_tokens:
|
||||
inputs += sum(self.counts.get(key, 0) for key in keys[2:])
|
||||
return {"input_tokens": inputs, "output_tokens": self.counts.get(self.output_key, 0)}
|
||||
|
||||
|
||||
class EventStreamingMixin:
|
||||
async def stream(self, request: ModelRequest) -> AsyncIterator[ModelEvent]:
|
||||
sequence = 0
|
||||
status = "completed"
|
||||
try:
|
||||
request, originals = prepare_tool_names(request)
|
||||
# Closing the public iterator must synchronously close every nested iterator.
|
||||
async with aclosing(self._events(request)) as events:
|
||||
async for kind, data in events:
|
||||
if kind == ModelEventType.tool_call_start and "name" in data:
|
||||
data = {**data, "name": originals.get(data["name"], data["name"])}
|
||||
if kind == ModelEventType.usage:
|
||||
data = {**data, "total_tokens": data["input_tokens"] + data["output_tokens"]}
|
||||
yield ModelEvent(event=kind, data=data, sequence=sequence,
|
||||
timestamp=datetime.now(timezone.utc))
|
||||
sequence += 1
|
||||
except ProviderError as exc:
|
||||
status = "failed"
|
||||
yield ModelEvent(event=ModelEventType.error, sequence=sequence,
|
||||
data={"code": exc.code, "message": exc.message},
|
||||
timestamp=datetime.now(timezone.utc))
|
||||
sequence += 1
|
||||
except (ValueError, TypeError, KeyError, IndexError, AttributeError, OverflowError):
|
||||
status = "failed"
|
||||
error = invalid_response()
|
||||
yield ModelEvent(event=ModelEventType.error, sequence=sequence,
|
||||
data={"code": error.code, "message": error.message},
|
||||
timestamp=datetime.now(timezone.utc))
|
||||
sequence += 1
|
||||
# CancelledError and GeneratorExit deliberately propagate without a Done event.
|
||||
yield ModelEvent(event=ModelEventType.done, sequence=sequence,
|
||||
data={"status": status},
|
||||
timestamp=datetime.now(timezone.utc))
|
||||
|
||||
|
||||
async def sse_objects(response: httpx.Response) -> AsyncIterator[dict]:
|
||||
"""Read SSE frames, accepting the adjacent data lines used by some gateways."""
|
||||
parts: list[str] = []
|
||||
event_name = ""
|
||||
|
||||
def decode() -> dict:
|
||||
value = "\n".join(parts)
|
||||
if value.strip() == "[DONE]":
|
||||
return {"type": "[DONE]"}
|
||||
try:
|
||||
data = object_value(json.loads(value))
|
||||
except (ValueError, TypeError) as exc:
|
||||
raise invalid_response() from exc
|
||||
if event_name and "type" not in data:
|
||||
data["type"] = event_name
|
||||
check_error(data)
|
||||
return data
|
||||
|
||||
async for line in response.aiter_lines():
|
||||
if not line:
|
||||
if parts:
|
||||
yield decode()
|
||||
parts = []
|
||||
event_name = ""
|
||||
elif line.startswith(":"):
|
||||
continue
|
||||
elif line.startswith("event:"):
|
||||
if parts:
|
||||
yield decode()
|
||||
parts = []
|
||||
event_name = line[6:].strip()
|
||||
elif line.startswith("data:"):
|
||||
if parts:
|
||||
# Legacy compatible endpoints sometimes omit blank separators.
|
||||
try:
|
||||
json.loads("\n".join(parts))
|
||||
except ValueError:
|
||||
pass
|
||||
else:
|
||||
yield decode()
|
||||
parts = []
|
||||
event_name = ""
|
||||
parts.append(line[5:].removeprefix(" "))
|
||||
if parts:
|
||||
yield decode()
|
||||
|
||||
|
||||
class HTTPProviderMixin:
|
||||
stream_path = "/chat/completions"
|
||||
stream_format = "sse"
|
||||
|
||||
def _headers(self) -> dict[str, str]:
|
||||
return {"Content-Type": "application/json"}
|
||||
|
||||
@staticmethod
|
||||
def _status_error(exc: httpx.HTTPStatusError) -> ProviderError:
|
||||
status = exc.response.status_code
|
||||
code = {400: "PROVIDER_INVALID_REQUEST", 401: "PROVIDER_AUTH_FAILED",
|
||||
403: "PROVIDER_AUTH_FAILED", 404: "MODEL_NOT_FOUND",
|
||||
408: "PROVIDER_TIMEOUT", 413: "PROVIDER_INVALID_REQUEST",
|
||||
422: "PROVIDER_INVALID_REQUEST", 429: "PROVIDER_RATE_LIMITED"}.get(
|
||||
status, "PROVIDER_UNAVAILABLE")
|
||||
return ProviderError(code, f"Provider returned HTTP {status}.")
|
||||
|
||||
async def _request(self, method: str, path: str, **kwargs) -> dict:
|
||||
headers = self._headers()
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=self.timeout_seconds, transport=self.transport) as client:
|
||||
response = await client.request(method, f"{self.base_url}{path}", headers=headers, **kwargs)
|
||||
response.raise_for_status()
|
||||
data = object_value(response.json())
|
||||
check_error(data)
|
||||
return data
|
||||
except httpx.TimeoutException as exc:
|
||||
raise ProviderError("PROVIDER_TIMEOUT", "Provider request timed out.") from exc
|
||||
except httpx.HTTPStatusError as exc:
|
||||
raise self._status_error(exc) from exc
|
||||
except httpx.HTTPError as exc:
|
||||
raise ProviderError("PROVIDER_UNAVAILABLE", "Provider is unavailable.") from exc
|
||||
except (ValueError, TypeError) as exc:
|
||||
raise invalid_response() from exc
|
||||
|
||||
async def _stream_json(self, payload: dict[str, object]) -> AsyncIterator[dict]:
|
||||
headers = self._headers()
|
||||
headers["Accept"] = "text/event-stream" if self.stream_format == "sse" else "application/x-ndjson"
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=self.timeout_seconds, transport=self.transport) as client:
|
||||
async with client.stream("POST", f"{self.base_url}{self.stream_path}",
|
||||
headers=headers, json=payload) as response:
|
||||
response.raise_for_status()
|
||||
if self.stream_format == "sse":
|
||||
async with aclosing(sse_objects(response)) as objects:
|
||||
async for data in objects:
|
||||
yield data
|
||||
else:
|
||||
async for line in response.aiter_lines():
|
||||
if line.strip():
|
||||
data = object_value(json.loads(line))
|
||||
check_error(data)
|
||||
yield data
|
||||
except httpx.TimeoutException as exc:
|
||||
raise ProviderError("PROVIDER_TIMEOUT", "Provider request timed out.") from exc
|
||||
except httpx.HTTPStatusError as exc:
|
||||
raise self._status_error(exc) from exc
|
||||
except httpx.HTTPError as exc:
|
||||
raise ProviderError("PROVIDER_UNAVAILABLE", "Provider is unavailable.") from exc
|
||||
except (ValueError, TypeError) as exc:
|
||||
raise invalid_response() from exc
|
||||
|
||||
@@ -1,16 +1,22 @@
|
||||
from uuid import uuid4
|
||||
import json
|
||||
from collections.abc import AsyncIterator
|
||||
from datetime import datetime, timezone
|
||||
from contextlib import aclosing
|
||||
from uuid import uuid4
|
||||
|
||||
import httpx
|
||||
|
||||
from app.contracts import ModelCapability, ModelEvent, ModelEventType, ModelInfo, ModelRequest
|
||||
from app.contracts import MessageRole, ModelCapability, ModelEventType, ModelInfo, ModelRequest
|
||||
from app.providers.base import ProviderError, ProviderToolCall, ProviderTurn
|
||||
from app.providers.http_base import TurnStreamingMixin, decode_tool_arguments
|
||||
from app.providers.tool_names import mapped_tool_names
|
||||
from app.providers.http_base import (
|
||||
EventStreamingMixin, HTTPProviderMixin, UsageTracker, decode_tool_arguments,
|
||||
invalid_response, list_value, object_value, string_value, truncated_stream,
|
||||
)
|
||||
|
||||
|
||||
class OllamaProvider(TurnStreamingMixin):
|
||||
class OllamaProvider(EventStreamingMixin, HTTPProviderMixin):
|
||||
stream_path = "/api/chat"
|
||||
stream_format = "jsonl"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str = "http://127.0.0.1:11434",
|
||||
@@ -21,126 +27,55 @@ class OllamaProvider(TurnStreamingMixin):
|
||||
self.timeout_seconds = timeout_seconds
|
||||
self.transport = transport
|
||||
|
||||
@mapped_tool_names
|
||||
async def complete(self, request: ModelRequest) -> ProviderTurn:
|
||||
messages = []
|
||||
if request.system:
|
||||
messages.append({"role": "system", "content": request.system})
|
||||
for message in request.messages:
|
||||
item: dict[str, object] = {
|
||||
"role": message.role.value,
|
||||
"content": message.content,
|
||||
}
|
||||
if message.tool_calls:
|
||||
item["tool_calls"] = [
|
||||
{
|
||||
"function": {
|
||||
"name": call.name,
|
||||
"arguments": call.arguments,
|
||||
}
|
||||
}
|
||||
for call in message.tool_calls
|
||||
]
|
||||
messages.append(item)
|
||||
payload: dict[str, object] = {
|
||||
"model": request.model,
|
||||
"messages": messages,
|
||||
"stream": False,
|
||||
}
|
||||
if request.tools:
|
||||
payload["tools"] = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tool.name,
|
||||
"description": tool.description,
|
||||
"parameters": tool.parameters,
|
||||
},
|
||||
}
|
||||
for tool in request.tools
|
||||
]
|
||||
data = await self._request("POST", "/api/chat", json=payload)
|
||||
message = data.get("message") or {}
|
||||
tool_calls = []
|
||||
for raw_call in message.get("tool_calls") or []:
|
||||
function = raw_call.get("function") or {}
|
||||
tool_calls.append(
|
||||
ProviderToolCall(
|
||||
tool_call_id=raw_call.get("id") or f"call_{uuid4().hex}",
|
||||
name=function.get("name") or "",
|
||||
data = await self._request("POST", self.stream_path, json=self._chat_payload(request, stream=False))
|
||||
message = object_value(data.get("message"))
|
||||
calls = [self._tool_call(raw) for raw in list_value(message.get("tool_calls", []))]
|
||||
content = message.get("content")
|
||||
if content is not None:
|
||||
content = string_value(content)
|
||||
return ProviderTurn(text=content or None, tool_calls=calls,
|
||||
**UsageTracker("prompt_eval_count", "eval_count").update(data))
|
||||
|
||||
@staticmethod
|
||||
def _tool_call(raw: object) -> ProviderToolCall:
|
||||
call = object_value(raw)
|
||||
function = object_value(call.get("function"))
|
||||
return ProviderToolCall(
|
||||
tool_call_id=string_value(call.get("id") or f"call_{uuid4().hex}"),
|
||||
name=string_value(function.get("name"), nonempty=True),
|
||||
arguments=decode_tool_arguments(function.get("arguments", {})),
|
||||
)
|
||||
)
|
||||
return ProviderTurn(
|
||||
text=message.get("content") or None,
|
||||
tool_calls=tool_calls,
|
||||
input_tokens=int(data.get("prompt_eval_count") or 0),
|
||||
output_tokens=int(data.get("eval_count") or 0),
|
||||
)
|
||||
|
||||
async def list_models(self) -> list[ModelInfo]:
|
||||
data = await self._request("GET", "/api/tags")
|
||||
return [
|
||||
ModelInfo(
|
||||
model=item["name"],
|
||||
display_name=item.get("name", ""),
|
||||
capabilities=[ModelCapability.chat, ModelCapability.streaming],
|
||||
)
|
||||
for item in data.get("models", [])
|
||||
if isinstance(item, dict) and item.get("name")
|
||||
]
|
||||
|
||||
async def stream(self, request: ModelRequest) -> AsyncIterator[ModelEvent]:
|
||||
payload = self._chat_payload(request, stream=True)
|
||||
sequence = 0
|
||||
|
||||
def event(kind: ModelEventType, data: dict | None = None) -> ModelEvent:
|
||||
nonlocal sequence
|
||||
item = ModelEvent(
|
||||
event=kind, sequence=sequence, data=data or {},
|
||||
timestamp=datetime.now(timezone.utc),
|
||||
)
|
||||
sequence += 1
|
||||
return item
|
||||
|
||||
try:
|
||||
async for data in self._stream_json(payload):
|
||||
message = data.get("message") or {}
|
||||
async def _events(self, request: ModelRequest):
|
||||
usage = UsageTracker("prompt_eval_count", "eval_count")
|
||||
async with aclosing(self._stream_json(self._chat_payload(request, stream=True))) as chunks:
|
||||
async for data in chunks:
|
||||
message = object_value(data.get("message", {}))
|
||||
if message.get("thinking"):
|
||||
yield event(ModelEventType.thinking_delta, {"text": message["thinking"]})
|
||||
yield ModelEventType.thinking_delta, {"text": string_value(message["thinking"])}
|
||||
if message.get("content"):
|
||||
yield event(ModelEventType.text_delta, {"text": message["content"]})
|
||||
for raw_call in message.get("tool_calls") or []:
|
||||
function = raw_call.get("function") or {}
|
||||
call_id = raw_call.get("id") or f"call_{uuid4().hex}"
|
||||
yield event(
|
||||
ModelEventType.tool_call_start,
|
||||
{"tool_call_id": call_id, "name": function.get("name") or ""},
|
||||
)
|
||||
yield event(
|
||||
ModelEventType.tool_call_delta,
|
||||
{
|
||||
"tool_call_id": call_id,
|
||||
"arguments_delta": json.dumps(
|
||||
function.get("arguments") or {}, ensure_ascii=False
|
||||
),
|
||||
},
|
||||
)
|
||||
yield event(ModelEventType.tool_call_end, {"tool_call_id": call_id})
|
||||
if data.get("done"):
|
||||
yield event(
|
||||
ModelEventType.usage,
|
||||
{
|
||||
"input_tokens": int(data.get("prompt_eval_count") or 0),
|
||||
"output_tokens": int(data.get("eval_count") or 0),
|
||||
},
|
||||
)
|
||||
yield event(ModelEventType.done)
|
||||
except ProviderError as exc:
|
||||
yield event(ModelEventType.error, {"code": exc.code, "message": exc.message})
|
||||
yield event(ModelEventType.done)
|
||||
yield ModelEventType.text_delta, {"text": string_value(message["content"])}
|
||||
for raw in list_value(message.get("tool_calls", [])):
|
||||
call = self._tool_call(raw)
|
||||
yield ModelEventType.tool_call_start, {"tool_call_id": call.tool_call_id, "name": call.name}
|
||||
yield ModelEventType.tool_call_delta, {
|
||||
"tool_call_id": call.tool_call_id,
|
||||
"arguments_delta": json.dumps(call.arguments, ensure_ascii=False),
|
||||
}
|
||||
yield ModelEventType.tool_call_end, {"tool_call_id": call.tool_call_id}
|
||||
if "done" in data and not isinstance(data["done"], bool):
|
||||
raise invalid_response()
|
||||
if "prompt_eval_count" in data or "eval_count" in data or data.get("done"):
|
||||
yield ModelEventType.usage, usage.update(data)
|
||||
if data.get("done") is True:
|
||||
return
|
||||
raise truncated_stream()
|
||||
|
||||
def _chat_payload(self, request: ModelRequest, *, stream: bool) -> dict[str, object]:
|
||||
messages = []
|
||||
names: dict[str, str] = {}
|
||||
if request.system:
|
||||
messages.append({"role": "system", "content": request.system})
|
||||
for message in request.messages:
|
||||
@@ -150,53 +85,49 @@ class OllamaProvider(TurnStreamingMixin):
|
||||
{"function": {"name": call.name, "arguments": call.arguments}}
|
||||
for call in message.tool_calls
|
||||
]
|
||||
names.update({call.tool_call_id: call.name for call in message.tool_calls})
|
||||
if message.role == MessageRole.tool:
|
||||
name = message.name or names.get(message.tool_call_id or "")
|
||||
if name:
|
||||
item["tool_name"] = name
|
||||
messages.append(item)
|
||||
payload: dict[str, object] = {
|
||||
"model": request.model, "messages": messages, "stream": stream
|
||||
"model": request.model, "messages": messages, "stream": stream,
|
||||
}
|
||||
if request.tools:
|
||||
payload["tools"] = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tool.name,
|
||||
"description": tool.description,
|
||||
"parameters": tool.parameters,
|
||||
},
|
||||
}
|
||||
for tool in request.tools
|
||||
{"type": "function", "function": {
|
||||
"name": tool.name, "description": tool.description, "parameters": tool.parameters,
|
||||
}} for tool in request.tools
|
||||
]
|
||||
options = {}
|
||||
if request.temperature is not None:
|
||||
options["temperature"] = request.temperature
|
||||
if request.max_tokens is not None:
|
||||
options["num_predict"] = request.max_tokens
|
||||
if options:
|
||||
payload["options"] = options
|
||||
if request.response_format:
|
||||
format_ = request.response_format
|
||||
if format_.get("type") == "json_object":
|
||||
payload["format"] = "json"
|
||||
elif format_.get("type") == "json_schema":
|
||||
payload["format"] = object_value(object_value(format_.get("json_schema")).get("schema"))
|
||||
else:
|
||||
payload["format"] = format_
|
||||
return payload
|
||||
|
||||
async def _stream_json(self, payload: dict[str, object]) -> AsyncIterator[dict]:
|
||||
try:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=self.timeout_seconds, transport=self.transport
|
||||
) as client:
|
||||
async with client.stream(
|
||||
"POST", f"{self.base_url}/api/chat", json=payload
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
async for line in response.aiter_lines():
|
||||
if not line.strip():
|
||||
continue
|
||||
try:
|
||||
data = json.loads(line)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ProviderError(
|
||||
"PROVIDER_INVALID_RESPONSE", "Ollama returned invalid JSONL."
|
||||
) from exc
|
||||
if isinstance(data, dict):
|
||||
yield data
|
||||
except httpx.TimeoutException as exc:
|
||||
raise ProviderError("PROVIDER_TIMEOUT", "Ollama request timed out.") from exc
|
||||
except httpx.HTTPStatusError as exc:
|
||||
raise ProviderError(
|
||||
"MODEL_NOT_FOUND" if exc.response.status_code == 404 else "PROVIDER_UNAVAILABLE",
|
||||
f"Ollama returned HTTP {exc.response.status_code}.",
|
||||
) from exc
|
||||
except httpx.HTTPError as exc:
|
||||
raise ProviderError("PROVIDER_UNAVAILABLE", "Ollama is unavailable.") from exc
|
||||
async def list_models(self) -> list[ModelInfo]:
|
||||
data = await self._request("GET", "/api/tags")
|
||||
return [
|
||||
ModelInfo(
|
||||
model=string_value(item["name"]), display_name=item["name"],
|
||||
capabilities=([ModelCapability.embedding] if "embed" in item["name"].lower()
|
||||
else [ModelCapability.chat, ModelCapability.streaming]),
|
||||
)
|
||||
for item in list_value(data.get("models"))
|
||||
if isinstance(item, dict) and isinstance(item.get("name"), str) and item["name"]
|
||||
]
|
||||
|
||||
async def test_connection(self, model: str | None = None) -> tuple[bool, str]:
|
||||
try:
|
||||
@@ -206,24 +137,3 @@ class OllamaProvider(TurnStreamingMixin):
|
||||
if model and model not in {item.model for item in models}:
|
||||
return False, f"Model is not installed: {model}"
|
||||
return True, f"Connected; discovered {len(models)} local model(s)."
|
||||
|
||||
async def _request(self, method: str, path: str, **kwargs) -> dict:
|
||||
try:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=self.timeout_seconds, transport=self.transport
|
||||
) as client:
|
||||
response = await client.request(method, f"{self.base_url}{path}", **kwargs)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
except httpx.TimeoutException as exc:
|
||||
raise ProviderError("PROVIDER_TIMEOUT", "Ollama request timed out.") from exc
|
||||
except httpx.HTTPStatusError as exc:
|
||||
raise ProviderError(
|
||||
"MODEL_NOT_FOUND" if exc.response.status_code == 404 else "PROVIDER_UNAVAILABLE",
|
||||
f"Ollama returned HTTP {exc.response.status_code}.",
|
||||
) from exc
|
||||
except (httpx.HTTPError, ValueError) as exc:
|
||||
raise ProviderError("PROVIDER_UNAVAILABLE", "Ollama is unavailable.") from exc
|
||||
if not isinstance(data, dict):
|
||||
raise ProviderError("PROVIDER_INVALID_RESPONSE", "Ollama returned non-object JSON.")
|
||||
return data
|
||||
|
||||
@@ -1,24 +1,20 @@
|
||||
import json
|
||||
from collections.abc import AsyncIterator
|
||||
from datetime import datetime, timezone
|
||||
from contextlib import aclosing
|
||||
from uuid import uuid4
|
||||
|
||||
import httpx
|
||||
|
||||
from app.contracts import (
|
||||
MessageRole,
|
||||
ModelCapability,
|
||||
ModelEvent,
|
||||
ModelEventType,
|
||||
ModelInfo,
|
||||
ModelRequest,
|
||||
)
|
||||
from app.contracts import MessageRole, ModelCapability, ModelEventType, ModelInfo, ModelRequest
|
||||
from app.providers.base import ProviderError, ProviderToolCall, ProviderTurn
|
||||
from app.providers.credentials import CredentialResolver, CredentialStoreError
|
||||
from app.providers.http_base import TurnStreamingMixin, decode_tool_arguments
|
||||
from app.providers.tool_names import mapped_tool_names
|
||||
from app.providers.http_base import (
|
||||
EventStreamingMixin, HTTPProviderMixin, UsageTracker, decode_tool_arguments,
|
||||
invalid_response, list_value, object_value, string_value, token_count, truncated_stream,
|
||||
)
|
||||
|
||||
|
||||
class OpenAICompatibleProvider(TurnStreamingMixin):
|
||||
class OpenAICompatibleProvider(EventStreamingMixin, HTTPProviderMixin):
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
@@ -33,50 +29,37 @@ class OpenAICompatibleProvider(TurnStreamingMixin):
|
||||
self.timeout_seconds = timeout_seconds
|
||||
self.transport = transport
|
||||
|
||||
@mapped_tool_names
|
||||
async def complete(self, request: ModelRequest) -> ProviderTurn:
|
||||
payload = self._payload(request, stream=False)
|
||||
|
||||
data = await self._request("POST", "/chat/completions", json=payload)
|
||||
try:
|
||||
message = data["choices"][0]["message"]
|
||||
except (KeyError, IndexError, TypeError) as exc:
|
||||
raise ProviderError("PROVIDER_INVALID_RESPONSE", "Missing completion message.") from exc
|
||||
|
||||
tool_calls = []
|
||||
for raw_call in message.get("tool_calls") or []:
|
||||
function = raw_call.get("function") or {}
|
||||
tool_calls.append(
|
||||
ProviderToolCall(
|
||||
tool_call_id=raw_call.get("id") or f"call_{uuid4().hex}",
|
||||
name=function.get("name") or "",
|
||||
data = await self._request("POST", self.stream_path, json=self._payload(request, stream=False))
|
||||
choices = list_value(data.get("choices"))
|
||||
if not choices:
|
||||
raise invalid_response()
|
||||
message = object_value(object_value(choices[0]).get("message"))
|
||||
calls = []
|
||||
for raw in list_value(message.get("tool_calls", [])):
|
||||
raw = object_value(raw)
|
||||
function = object_value(raw.get("function"))
|
||||
calls.append(ProviderToolCall(
|
||||
tool_call_id=string_value(raw.get("id") or f"call_{uuid4().hex}"),
|
||||
name=string_value(function.get("name"), nonempty=True),
|
||||
arguments=decode_tool_arguments(function.get("arguments", "{}")),
|
||||
)
|
||||
)
|
||||
usage = data.get("usage") or {}
|
||||
return ProviderTurn(
|
||||
text=message.get("content"),
|
||||
tool_calls=tool_calls,
|
||||
input_tokens=int(usage.get("prompt_tokens") or 0),
|
||||
output_tokens=int(usage.get("completion_tokens") or 0),
|
||||
)
|
||||
))
|
||||
text = message.get("content")
|
||||
if text is not None:
|
||||
text = string_value(text)
|
||||
usage = UsageTracker("prompt_tokens", "completion_tokens").update(data.get("usage") or {})
|
||||
return ProviderTurn(text=text, tool_calls=calls, **usage)
|
||||
|
||||
def _payload(self, request: ModelRequest, *, stream: bool) -> dict[str, object]:
|
||||
payload: dict[str, object] = {
|
||||
"model": request.model,
|
||||
"messages": self._messages(request),
|
||||
"stream": stream,
|
||||
"model": request.model, "messages": self._messages(request), "stream": stream,
|
||||
}
|
||||
if request.tools:
|
||||
payload["tools"] = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tool.name,
|
||||
"description": tool.description,
|
||||
"parameters": tool.parameters,
|
||||
},
|
||||
}
|
||||
for tool in request.tools
|
||||
{"type": "function", "function": {
|
||||
"name": tool.name, "description": tool.description, "parameters": tool.parameters,
|
||||
}} for tool in request.tools
|
||||
]
|
||||
if request.temperature is not None:
|
||||
payload["temperature"] = request.temperature
|
||||
@@ -84,124 +67,78 @@ class OpenAICompatibleProvider(TurnStreamingMixin):
|
||||
payload["max_tokens"] = request.max_tokens
|
||||
if request.response_format is not None:
|
||||
payload["response_format"] = request.response_format
|
||||
|
||||
if stream:
|
||||
payload["stream_options"] = {"include_usage": True}
|
||||
return payload
|
||||
|
||||
async def stream(self, request: ModelRequest) -> AsyncIterator[ModelEvent]:
|
||||
sequence = 0
|
||||
open_calls: dict[int, str] = {}
|
||||
|
||||
def event(kind: ModelEventType, data: dict | None = None) -> ModelEvent:
|
||||
nonlocal sequence
|
||||
item = ModelEvent(
|
||||
event=kind,
|
||||
sequence=sequence,
|
||||
data=data or {},
|
||||
timestamp=datetime.now(timezone.utc),
|
||||
)
|
||||
sequence += 1
|
||||
return item
|
||||
|
||||
try:
|
||||
async for data in self._stream_json(self._payload(request, stream=True)):
|
||||
usage = data.get("usage") or {}
|
||||
if usage:
|
||||
yield event(
|
||||
ModelEventType.usage,
|
||||
{
|
||||
"input_tokens": int(usage.get("prompt_tokens") or 0),
|
||||
"output_tokens": int(usage.get("completion_tokens") or 0),
|
||||
},
|
||||
)
|
||||
choices = data.get("choices") or []
|
||||
async def _events(self, request: ModelRequest):
|
||||
calls: dict[int, dict] = {}
|
||||
usage = UsageTracker("prompt_tokens", "completion_tokens")
|
||||
finished = False
|
||||
seen = False
|
||||
async with aclosing(self._stream_json(self._payload(request, stream=True))) as chunks:
|
||||
async for data in chunks:
|
||||
if data.get("type") == "[DONE]":
|
||||
if not seen:
|
||||
raise invalid_response()
|
||||
finished = True
|
||||
break
|
||||
if data.get("usage") is not None:
|
||||
yield ModelEventType.usage, usage.update(data["usage"])
|
||||
choices = list_value(data.get("choices", []))
|
||||
if not choices:
|
||||
continue
|
||||
choice = choices[0]
|
||||
delta = choice.get("delta") or {}
|
||||
seen = True
|
||||
choice = object_value(choices[0])
|
||||
delta = object_value(choice.get("delta") or {})
|
||||
if delta.get("reasoning_content"):
|
||||
yield event(
|
||||
ModelEventType.thinking_delta,
|
||||
{"text": delta["reasoning_content"]},
|
||||
)
|
||||
yield ModelEventType.thinking_delta, {"text": string_value(delta["reasoning_content"])}
|
||||
if delta.get("content"):
|
||||
yield event(ModelEventType.text_delta, {"text": delta["content"]})
|
||||
for raw_call in delta.get("tool_calls") or []:
|
||||
index = int(raw_call.get("index") or 0)
|
||||
function = raw_call.get("function") or {}
|
||||
call_id = raw_call.get("id") or open_calls.get(index) or f"call_{uuid4().hex}"
|
||||
if index not in open_calls:
|
||||
open_calls[index] = call_id
|
||||
yield event(
|
||||
ModelEventType.tool_call_start,
|
||||
{"tool_call_id": call_id, "name": function.get("name") or ""},
|
||||
)
|
||||
if function.get("arguments"):
|
||||
yield event(
|
||||
ModelEventType.tool_call_delta,
|
||||
{
|
||||
"tool_call_id": open_calls[index],
|
||||
"arguments_delta": function["arguments"],
|
||||
},
|
||||
)
|
||||
if choice.get("finish_reason") == "tool_calls":
|
||||
for call_id in open_calls.values():
|
||||
yield event(
|
||||
ModelEventType.tool_call_end, {"tool_call_id": call_id}
|
||||
)
|
||||
open_calls.clear()
|
||||
for call_id in open_calls.values():
|
||||
yield event(ModelEventType.tool_call_end, {"tool_call_id": call_id})
|
||||
yield event(ModelEventType.done)
|
||||
except ProviderError as exc:
|
||||
yield event(ModelEventType.error, {"code": exc.code, "message": exc.message})
|
||||
yield event(ModelEventType.done)
|
||||
|
||||
async def _stream_json(self, payload: dict[str, object]) -> AsyncIterator[dict]:
|
||||
headers = self._headers()
|
||||
try:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=self.timeout_seconds, transport=self.transport
|
||||
) as client:
|
||||
async with client.stream(
|
||||
"POST", f"{self.base_url}/chat/completions", headers=headers, json=payload
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
async for line in response.aiter_lines():
|
||||
if not line.startswith("data:"):
|
||||
continue
|
||||
value = line[5:].strip()
|
||||
if not value or value == "[DONE]":
|
||||
continue
|
||||
try:
|
||||
data = json.loads(value)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ProviderError(
|
||||
"PROVIDER_INVALID_RESPONSE", "Provider returned invalid SSE JSON."
|
||||
) from exc
|
||||
if isinstance(data, dict):
|
||||
yield data
|
||||
except httpx.TimeoutException as exc:
|
||||
raise ProviderError("PROVIDER_TIMEOUT", "Provider request timed out.") from exc
|
||||
except httpx.HTTPStatusError as exc:
|
||||
raise self._status_error(exc) from exc
|
||||
except httpx.HTTPError as exc:
|
||||
raise ProviderError("PROVIDER_UNAVAILABLE", "Provider is unavailable.") from exc
|
||||
yield ModelEventType.text_delta, {"text": string_value(delta["content"])}
|
||||
for raw in list_value(delta.get("tool_calls", [])):
|
||||
raw = object_value(raw)
|
||||
index = token_count(raw.get("index", 0))
|
||||
function = object_value(raw.get("function") or {})
|
||||
call = calls.setdefault(index, {"id": "", "name": "", "arguments": ""})
|
||||
if raw.get("id"):
|
||||
call["id"] = string_value(raw["id"])
|
||||
if function.get("name"):
|
||||
call["name"] += string_value(function["name"])
|
||||
fragment = string_value(function.get("arguments", ""))
|
||||
call["arguments"] += fragment
|
||||
if choice.get("finish_reason"):
|
||||
finished = True
|
||||
if not finished:
|
||||
raise truncated_stream()
|
||||
for call in calls.values():
|
||||
if not call["name"]:
|
||||
raise invalid_response()
|
||||
decode_tool_arguments(call["arguments"] or "{}")
|
||||
# A name can span multiple chunks; publish only the complete identity.
|
||||
call["id"] = call["id"] or f"call_{uuid4().hex}"
|
||||
yield ModelEventType.tool_call_start, {"tool_call_id": call["id"], "name": call["name"]}
|
||||
yield ModelEventType.tool_call_delta, {"tool_call_id": call["id"], "arguments_delta": call["arguments"] or "{}"}
|
||||
yield ModelEventType.tool_call_end, {"tool_call_id": call["id"]}
|
||||
|
||||
async def list_models(self) -> list[ModelInfo]:
|
||||
data = await self._request("GET", "/models")
|
||||
return [
|
||||
ModelInfo(
|
||||
model=item["id"],
|
||||
display_name=item["id"],
|
||||
capabilities=[
|
||||
ModelCapability.chat,
|
||||
ModelCapability.tool_calling,
|
||||
ModelCapability.streaming,
|
||||
],
|
||||
)
|
||||
for item in data.get("data", [])
|
||||
if isinstance(item, dict) and item.get("id")
|
||||
]
|
||||
return [ModelInfo(model=string_value(item["id"]), display_name=item["id"],
|
||||
capabilities=self._model_capabilities(string_value(item["id"])))
|
||||
for item in list_value(data.get("data"))
|
||||
if isinstance(item, dict) and item.get("id")]
|
||||
|
||||
@staticmethod
|
||||
def _model_capabilities(model: str) -> list[ModelCapability]:
|
||||
# /models does not advertise capabilities. Avoid known non-chat families;
|
||||
# these are discovery hints, not a guarantee of support by a gateway.
|
||||
name = model.lower()
|
||||
if "embed" in name or name.startswith(("bge-", "bge/")):
|
||||
return [ModelCapability.embedding]
|
||||
if any(marker in name for marker in (
|
||||
"whisper", "tts", "transcri", "audio", "realtime", "dall-e", "image", "moderation", "rerank",
|
||||
)):
|
||||
return []
|
||||
return [ModelCapability.chat]
|
||||
|
||||
async def test_connection(self, model: str | None = None) -> tuple[bool, str]:
|
||||
try:
|
||||
@@ -217,73 +154,30 @@ class OpenAICompatibleProvider(TurnStreamingMixin):
|
||||
if request.system:
|
||||
result.append({"role": "system", "content": request.system})
|
||||
for message in request.messages:
|
||||
item: dict[str, object] = {
|
||||
"role": message.role.value,
|
||||
"content": message.content,
|
||||
}
|
||||
item: dict[str, object] = {"role": message.role.value, "content": message.content}
|
||||
if message.name:
|
||||
item["name"] = message.name
|
||||
if message.role == MessageRole.tool and message.tool_call_id:
|
||||
item["tool_call_id"] = message.tool_call_id
|
||||
if message.tool_calls:
|
||||
item["tool_calls"] = [
|
||||
{
|
||||
"id": call.tool_call_id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": call.name,
|
||||
"arguments": json.dumps(call.arguments),
|
||||
},
|
||||
}
|
||||
for call in message.tool_calls
|
||||
{"id": call.tool_call_id, "type": "function", "function": {
|
||||
"name": call.name, "arguments": json.dumps(call.arguments),
|
||||
}} for call in message.tool_calls
|
||||
]
|
||||
result.append(item)
|
||||
return result
|
||||
|
||||
async def _request(self, method: str, path: str, **kwargs) -> dict:
|
||||
headers = self._headers()
|
||||
try:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=self.timeout_seconds, transport=self.transport
|
||||
) as client:
|
||||
response = await client.request(
|
||||
method, f"{self.base_url}{path}", headers=headers, **kwargs
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
except httpx.TimeoutException as exc:
|
||||
raise ProviderError("PROVIDER_TIMEOUT", "Provider request timed out.") from exc
|
||||
except httpx.HTTPStatusError as exc:
|
||||
raise self._status_error(exc) from exc
|
||||
except (httpx.HTTPError, ValueError) as exc:
|
||||
raise ProviderError("PROVIDER_UNAVAILABLE", "Provider is unavailable.") from exc
|
||||
if not isinstance(data, dict):
|
||||
raise ProviderError("PROVIDER_INVALID_RESPONSE", "Provider returned non-object JSON.")
|
||||
return data
|
||||
|
||||
def _headers(self) -> dict[str, str]:
|
||||
headers = {"Content-Type": "application/json"}
|
||||
try:
|
||||
api_key = self.credentials.resolve(self.credential_id)
|
||||
except CredentialStoreError as exc:
|
||||
raise ProviderError(
|
||||
"PROVIDER_CREDENTIAL_UNAVAILABLE",
|
||||
"Credential could not be decrypted by the AI Core.",
|
||||
) from exc
|
||||
raise ProviderError("PROVIDER_CREDENTIAL_UNAVAILABLE",
|
||||
"Credential could not be decrypted by the AI Core.") from exc
|
||||
if self.credential_id and not api_key:
|
||||
raise ProviderError(
|
||||
"PROVIDER_CREDENTIAL_MISSING",
|
||||
f'Credential "{self.credential_id}" is not available in the AI Core process.',
|
||||
)
|
||||
raise ProviderError("PROVIDER_CREDENTIAL_MISSING",
|
||||
"Credential is not available in the AI Core process.")
|
||||
if api_key:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
return headers
|
||||
|
||||
@staticmethod
|
||||
def _status_error(exc: httpx.HTTPStatusError) -> ProviderError:
|
||||
code = {
|
||||
401: "PROVIDER_AUTH_FAILED",
|
||||
404: "MODEL_NOT_FOUND",
|
||||
429: "PROVIDER_RATE_LIMITED",
|
||||
}.get(exc.response.status_code, "PROVIDER_UNAVAILABLE")
|
||||
return ProviderError(code, f"Provider returned HTTP {exc.response.status_code}.")
|
||||
|
||||
@@ -0,0 +1,168 @@
|
||||
"""Native /responses adapter; stateless history uses function_call/output items."""
|
||||
|
||||
import json
|
||||
from contextlib import aclosing
|
||||
|
||||
from app.contracts import MessageRole, ModelEventType, ModelRequest
|
||||
from app.providers.base import ProviderError, ProviderToolCall, ProviderTurn
|
||||
from app.providers.http_base import (
|
||||
UsageTracker, check_error, decode_tool_arguments, invalid_response, list_value,
|
||||
object_value, remote_error, string_value, token_count, truncated_stream,
|
||||
)
|
||||
from app.providers.openai_compatible import OpenAICompatibleProvider
|
||||
from app.providers.tool_names import mapped_tool_names
|
||||
|
||||
|
||||
class OpenAIResponsesProvider(OpenAICompatibleProvider):
|
||||
stream_path = "/responses"
|
||||
|
||||
def _payload(self, request: ModelRequest, *, stream: bool) -> dict[str, object]:
|
||||
inputs = []
|
||||
for message in request.messages:
|
||||
if message.role == MessageRole.tool:
|
||||
if not message.tool_call_id:
|
||||
raise ProviderError("PROVIDER_INVALID_REQUEST", "Tool result requires a call identifier.")
|
||||
inputs.append({"type": "function_call_output", "call_id": message.tool_call_id,
|
||||
"output": message.content})
|
||||
continue
|
||||
if message.content or not message.tool_calls:
|
||||
inputs.append({"role": message.role.value, "content": message.content})
|
||||
for call in message.tool_calls:
|
||||
inputs.append({"type": "function_call", "call_id": call.tool_call_id,
|
||||
"name": call.name, "arguments": json.dumps(call.arguments)})
|
||||
payload: dict[str, object] = {"model": request.model, "input": inputs, "stream": stream}
|
||||
if request.system:
|
||||
payload["instructions"] = request.system
|
||||
if request.tools:
|
||||
payload["tools"] = [{"type": "function", "name": tool.name,
|
||||
"description": tool.description, "parameters": tool.parameters}
|
||||
for tool in request.tools]
|
||||
if request.temperature is not None:
|
||||
payload["temperature"] = request.temperature
|
||||
if request.max_tokens is not None:
|
||||
payload["max_output_tokens"] = request.max_tokens
|
||||
if request.response_format is not None:
|
||||
format_ = dict(request.response_format)
|
||||
if format_.get("type") == "json_schema":
|
||||
format_ = {"type": "json_schema", **object_value(format_.get("json_schema"))}
|
||||
payload["text"] = {"format": format_}
|
||||
return payload
|
||||
|
||||
@staticmethod
|
||||
def _check_response(data: dict) -> None:
|
||||
check_error(data)
|
||||
status = data.get("status")
|
||||
if status == "incomplete":
|
||||
raise ProviderError("PROVIDER_INCOMPLETE_RESPONSE", "Provider response is incomplete.")
|
||||
if status == "failed":
|
||||
raise remote_error(data.get("error"))
|
||||
if status is not None and status != "completed":
|
||||
raise invalid_response()
|
||||
|
||||
@mapped_tool_names
|
||||
async def complete(self, request: ModelRequest) -> ProviderTurn:
|
||||
data = await self._request("POST", self.stream_path, json=self._payload(request, stream=False))
|
||||
self._check_response(data)
|
||||
texts = []
|
||||
calls = []
|
||||
for raw in list_value(data.get("output")):
|
||||
item = object_value(raw)
|
||||
if item.get("type") == "message":
|
||||
for raw_part in list_value(item.get("content")):
|
||||
part = object_value(raw_part)
|
||||
if part.get("type") == "output_text":
|
||||
texts.append(string_value(part.get("text")))
|
||||
elif part.get("type") == "refusal":
|
||||
texts.append(string_value(part.get("refusal")))
|
||||
elif item.get("type") == "function_call":
|
||||
calls.append(ProviderToolCall(
|
||||
tool_call_id=string_value(item.get("call_id"), nonempty=True),
|
||||
name=string_value(item.get("name"), nonempty=True),
|
||||
arguments=decode_tool_arguments(item.get("arguments")),
|
||||
))
|
||||
return ProviderTurn(text="".join(texts) or None, tool_calls=calls,
|
||||
**UsageTracker().update(data.get("usage") or {}))
|
||||
|
||||
async def _events(self, request: ModelRequest):
|
||||
calls: dict[int, dict] = {}
|
||||
usage = UsageTracker()
|
||||
|
||||
def finish_call(index: int, final: object = None):
|
||||
call = calls[index]
|
||||
if call["ended"]:
|
||||
return []
|
||||
events = []
|
||||
if final is not None:
|
||||
arguments = string_value(final)
|
||||
if not arguments.startswith(call["arguments"]):
|
||||
raise invalid_response()
|
||||
remainder = arguments[len(call["arguments"]):]
|
||||
if remainder:
|
||||
events.append((ModelEventType.tool_call_delta,
|
||||
{"tool_call_id": call["id"], "arguments_delta": remainder}))
|
||||
call["arguments"] = arguments
|
||||
decode_tool_arguments(call["arguments"])
|
||||
call["ended"] = True
|
||||
events.append((ModelEventType.tool_call_end, {"tool_call_id": call["id"]}))
|
||||
return events
|
||||
|
||||
async with aclosing(self._stream_json(self._payload(request, stream=True))) as chunks:
|
||||
async for data in chunks:
|
||||
kind = string_value(data.get("type"), nonempty=True)
|
||||
if kind in {"response.failed", "response.incomplete"}:
|
||||
response = object_value(data.get("response"))
|
||||
self._check_response({**response, "status": kind.split(".")[1]})
|
||||
elif kind in {"response.output_text.delta", "response.refusal.delta"}:
|
||||
yield ModelEventType.text_delta, {"text": string_value(data.get("delta"))}
|
||||
elif kind in {"response.reasoning_summary_text.delta", "response.reasoning_text.delta"}:
|
||||
yield ModelEventType.thinking_delta, {"text": string_value(data.get("delta"))}
|
||||
elif kind in {"response.output_item.added", "response.output_item.done"}:
|
||||
item = object_value(data.get("item"))
|
||||
if item.get("type") != "function_call":
|
||||
continue
|
||||
index = token_count(data.get("output_index"))
|
||||
call_id = string_value(item.get("call_id"), nonempty=True)
|
||||
name = string_value(item.get("name"), nonempty=True)
|
||||
if index not in calls:
|
||||
calls[index] = {"id": call_id, "name": name, "arguments": "", "ended": False,
|
||||
"item_id": item.get("id")}
|
||||
yield ModelEventType.tool_call_start, {"tool_call_id": call_id, "name": name}
|
||||
elif calls[index]["id"] != call_id or calls[index]["name"] != name:
|
||||
raise invalid_response()
|
||||
if kind == "response.output_item.done":
|
||||
for event in finish_call(index, item.get("arguments")):
|
||||
yield event
|
||||
elif item.get("arguments"):
|
||||
arguments = string_value(item["arguments"])
|
||||
calls[index]["arguments"] += arguments
|
||||
yield ModelEventType.tool_call_delta, {"tool_call_id": call_id, "arguments_delta": arguments}
|
||||
elif kind in {"response.function_call_arguments.delta", "response.function_call_arguments.done"}:
|
||||
index = token_count(data.get("output_index"))
|
||||
call = calls.get(index)
|
||||
if call is None or (data.get("item_id") and call["item_id"] != data["item_id"]):
|
||||
raise invalid_response()
|
||||
if kind.endswith(".done"):
|
||||
for event in finish_call(index, data.get("arguments")):
|
||||
yield event
|
||||
else:
|
||||
if call["ended"]:
|
||||
raise invalid_response()
|
||||
fragment = string_value(data.get("delta"))
|
||||
call["arguments"] += fragment
|
||||
yield ModelEventType.tool_call_delta, {"tool_call_id": call["id"], "arguments_delta": fragment}
|
||||
elif kind == "response.completed":
|
||||
response = object_value(data.get("response"))
|
||||
self._check_response(response)
|
||||
if any(not call["ended"] for call in calls.values()):
|
||||
raise truncated_stream()
|
||||
if response.get("usage") is not None:
|
||||
yield ModelEventType.usage, usage.update(response["usage"])
|
||||
return
|
||||
elif kind == "[DONE]":
|
||||
raise truncated_stream()
|
||||
elif kind in {"response.created", "response.in_progress"}:
|
||||
response = object_value(data.get("response"))
|
||||
check_error(response)
|
||||
if response.get("usage") is not None:
|
||||
yield ModelEventType.usage, usage.update(response["usage"])
|
||||
raise truncated_stream()
|
||||
@@ -1,5 +1,10 @@
|
||||
from dataclasses import dataclass
|
||||
from time import perf_counter
|
||||
from pathlib import Path
|
||||
|
||||
from app.config import get_settings
|
||||
from app.database.db import connect
|
||||
from app.errors import ApiError
|
||||
|
||||
from app.contracts import ModelInfo, ProviderConfig, ProviderTestResponse
|
||||
from app.providers.base import ModelProvider
|
||||
@@ -16,20 +21,64 @@ class RegisteredProvider:
|
||||
|
||||
|
||||
class ProviderRegistry:
|
||||
def __init__(self) -> None:
|
||||
def __init__(self, factory=None) -> None:
|
||||
self._providers: dict[str, RegisteredProvider] = {}
|
||||
self._factory = factory
|
||||
self._loaded_path: Path | None = None
|
||||
|
||||
def _restore(self) -> None:
|
||||
if self._factory is None or self._loaded_path == get_settings().db_path:
|
||||
return
|
||||
conn = connect()
|
||||
try:
|
||||
conn.execute("CREATE TABLE IF NOT EXISTS provider_configs (provider_id TEXT PRIMARY KEY, config_json TEXT NOT NULL)")
|
||||
restored = {}
|
||||
for row in conn.execute("SELECT config_json FROM provider_configs"):
|
||||
config = ProviderConfig.model_validate_json(row["config_json"])
|
||||
if config.provider_id == "mock":
|
||||
raise ValueError("reserved provider")
|
||||
restored[config.provider_id] = RegisteredProvider(config, self._factory.build(config))
|
||||
if "mock" in self._providers:
|
||||
restored["mock"] = self._providers["mock"]
|
||||
self._providers = restored
|
||||
self._loaded_path = get_settings().db_path
|
||||
except (ValueError, TypeError) as exc:
|
||||
raise ApiError(500, "PROVIDER_STORAGE_INVALID", "Saved provider configuration could not be loaded.") from exc
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
def _save(self, config: ProviderConfig) -> None:
|
||||
if self._factory is None or config.provider_id == "mock":
|
||||
return
|
||||
conn = connect()
|
||||
try:
|
||||
conn.execute("INSERT OR REPLACE INTO provider_configs VALUES (?, ?)", (config.provider_id, config.model_dump_json()))
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
def register(self, config: ProviderConfig, adapter: ModelProvider) -> None:
|
||||
if config.provider_id != "mock":
|
||||
self._restore()
|
||||
if config.provider_id in self._providers:
|
||||
raise ValueError(f"Provider already registered: {config.provider_id}")
|
||||
self._save(config)
|
||||
self._providers[config.provider_id] = RegisteredProvider(config=config, adapter=adapter)
|
||||
|
||||
def unregister(self, provider_id: str) -> None:
|
||||
self._restore()
|
||||
if self._factory is not None:
|
||||
conn = connect()
|
||||
try:
|
||||
conn.execute("DELETE FROM provider_configs WHERE provider_id = ?", (provider_id,))
|
||||
finally:
|
||||
conn.close()
|
||||
self._providers.pop(provider_id, None)
|
||||
|
||||
def replace(self, config: ProviderConfig, adapter: ModelProvider) -> None:
|
||||
self._restore()
|
||||
if config.provider_id not in self._providers:
|
||||
raise ProviderNotFoundError(config.provider_id)
|
||||
self._save(config)
|
||||
self._providers[config.provider_id] = RegisteredProvider(config=config, adapter=adapter)
|
||||
|
||||
def get(self, provider_id: str) -> RegisteredProvider:
|
||||
@@ -39,12 +88,14 @@ class ProviderRegistry:
|
||||
return provider
|
||||
|
||||
def get_any(self, provider_id: str) -> RegisteredProvider:
|
||||
self._restore()
|
||||
try:
|
||||
return self._providers[provider_id]
|
||||
except KeyError as exc:
|
||||
raise ProviderNotFoundError(provider_id) from exc
|
||||
|
||||
def list_configs(self) -> list[ProviderConfig]:
|
||||
self._restore()
|
||||
return [item.config.model_copy(deep=True) for item in self._providers.values()]
|
||||
|
||||
async def list_models(self, provider_id: str) -> list[ModelInfo]:
|
||||
|
||||
@@ -0,0 +1,291 @@
|
||||
"""Capability routing: validated remote results, then an explicit local backend.
|
||||
|
||||
Phase E supplies HTTP adapters and injectable local contracts. Hash embeddings are
|
||||
still a development placeholder; speech models are installed in phase F.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Protocol
|
||||
|
||||
import httpx
|
||||
|
||||
from app.contracts import (
|
||||
EmbeddingResult, LocalBackendStatus, ModelBinding, ModelRoutingConfig,
|
||||
ModelRoutingResponse, ProviderType, SpeakerMatchResult,
|
||||
)
|
||||
from app.database.db import connect, transaction
|
||||
from app.errors import ApiError
|
||||
from app.providers.base import ProviderError
|
||||
from app.providers.credentials import CredentialResolver, CredentialStoreError
|
||||
from app.providers.registry import ProviderNotFoundError, ProviderRegistry
|
||||
from app.retrieval.embedding import EmbeddingProvider, HashEmbeddingProvider
|
||||
from app.retrieval.provenance import record_embedding
|
||||
|
||||
CAPABILITIES = ("embedding", "transcription", "speaker_matching")
|
||||
HTTP_TYPES = {ProviderType.openai_chat, ProviderType.openai_compatible}
|
||||
MAX_MEDIA_BYTES = 25 * 1024 * 1024
|
||||
MAX_RESPONSE_BYTES = 16 * 1024 * 1024
|
||||
|
||||
|
||||
class LocalSpeechBackend(Protocol):
|
||||
available: bool
|
||||
|
||||
async def transcribe(self, source: Path, language: str | None) -> str: ...
|
||||
|
||||
async def match(self, source: Path, reference: Path) -> float: ...
|
||||
|
||||
|
||||
class PendingSpeechBackend:
|
||||
available = False
|
||||
|
||||
async def transcribe(self, source: Path, language: str | None) -> str:
|
||||
raise ProviderError("LOCAL_MODEL_NOT_INSTALLED", "本地音频转写模型尚未安装,将在阶段 F 接入。")
|
||||
|
||||
async def match(self, source: Path, reference: Path) -> float:
|
||||
raise ProviderError("LOCAL_MODEL_NOT_INSTALLED", "本地声纹模型尚未安装,将在阶段 F 接入。")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RoutedTranscript:
|
||||
text: str
|
||||
source: str
|
||||
fallback_reason: str | None = None
|
||||
|
||||
|
||||
def invalid_response() -> ProviderError:
|
||||
return ProviderError("PROVIDER_INVALID_RESPONSE", "Model API returned an invalid result.")
|
||||
|
||||
|
||||
def finite_number(value: object) -> bool:
|
||||
if type(value) not in (int, float):
|
||||
return False
|
||||
try:
|
||||
return math.isfinite(value)
|
||||
except (OverflowError, ValueError):
|
||||
return False
|
||||
|
||||
|
||||
class ModelRoutingService:
|
||||
def __init__(self, providers: ProviderRegistry, credentials: CredentialResolver, *,
|
||||
local_embedding: EmbeddingProvider | None = None,
|
||||
local_speech: LocalSpeechBackend | None = None,
|
||||
transport: httpx.AsyncBaseTransport | None = None) -> None:
|
||||
self.providers = providers
|
||||
self.credentials = credentials
|
||||
self.local_embedding = local_embedding or HashEmbeddingProvider()
|
||||
self.local_speech = local_speech or PendingSpeechBackend()
|
||||
self.transport = transport
|
||||
|
||||
@staticmethod
|
||||
def _connection():
|
||||
conn = connect()
|
||||
conn.execute("CREATE TABLE IF NOT EXISTS model_routing (id INTEGER PRIMARY KEY CHECK(id=1), config_json TEXT NOT NULL)")
|
||||
return conn
|
||||
|
||||
def configuration(self) -> ModelRoutingConfig:
|
||||
conn = self._connection()
|
||||
try:
|
||||
row = conn.execute("SELECT config_json FROM model_routing WHERE id=1").fetchone()
|
||||
return ModelRoutingConfig.model_validate_json(row[0]) if row else ModelRoutingConfig()
|
||||
except ValueError as exc:
|
||||
raise ApiError(500, "MODEL_ROUTING_STORAGE_INVALID", "Saved model routing could not be loaded.") from exc
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
def describe(self) -> ModelRoutingResponse:
|
||||
return ModelRoutingResponse(config=self.configuration(), local_backends=[
|
||||
LocalBackendStatus(capability="embedding", status="placeholder" if isinstance(self.local_embedding, HashEmbeddingProvider) else "ready",
|
||||
message="当前为 hash-v1 确定性占位向量,真实本地语义模型尚未集成。" if isinstance(self.local_embedding, HashEmbeddingProvider) else "本地 Embedding 模型已就绪。"),
|
||||
*[LocalBackendStatus(capability=capability, status="ready" if self.local_speech.available else "not_installed",
|
||||
message="本地模型已就绪。" if self.local_speech.available else "阶段 F 接入本地模型;当前保留回退接口。")
|
||||
for capability in ("transcription", "speaker_matching")],
|
||||
])
|
||||
|
||||
def update(self, config: ModelRoutingConfig) -> ModelRoutingResponse:
|
||||
for capability in CAPABILITIES:
|
||||
binding = getattr(config, capability)
|
||||
if binding:
|
||||
try:
|
||||
provider = self.providers.get_any(binding.provider_id).config
|
||||
except ProviderNotFoundError as exc:
|
||||
raise ApiError(422, "PROVIDER_NOT_FOUND", "请选择已保存的提供商。") from exc
|
||||
if provider.provider_type not in HTTP_TYPES:
|
||||
raise ApiError(422, "MODEL_ROUTING_PROTOCOL_UNSUPPORTED", "该能力当前需要 OpenAI Compatible HTTP 接口。")
|
||||
conn = self._connection()
|
||||
try:
|
||||
with transaction(conn):
|
||||
row = conn.execute("SELECT config_json FROM model_routing WHERE id=1").fetchone()
|
||||
current = ModelRoutingConfig.model_validate_json(row[0]) if row else ModelRoutingConfig()
|
||||
if current.version != config.version:
|
||||
raise ApiError(409, "MODEL_ROUTING_VERSION_CONFLICT", "配置已更新,请重新加载后再保存。")
|
||||
saved = config.model_copy(update={"version": config.version + 1})
|
||||
conn.execute("INSERT OR REPLACE INTO model_routing VALUES (1, ?)", (saved.model_dump_json(),))
|
||||
finally:
|
||||
conn.close()
|
||||
return self.describe()
|
||||
|
||||
def uses_provider(self, provider_id: str) -> bool:
|
||||
config = self.configuration()
|
||||
return any(binding and binding.provider_id == provider_id for binding in
|
||||
(getattr(config, name) for name in CAPABILITIES))
|
||||
|
||||
def _remote(self, binding: ModelBinding) -> tuple[str, dict[str, str]]:
|
||||
try:
|
||||
provider = self.providers.get(binding.provider_id).config
|
||||
except ProviderNotFoundError as exc:
|
||||
raise ProviderError("PROVIDER_UNAVAILABLE", "Configured provider is unavailable.") from exc
|
||||
if provider.provider_type not in HTTP_TYPES:
|
||||
raise ProviderError("PROVIDER_CAPABILITY_UNSUPPORTED", "Provider does not support this HTTP capability.")
|
||||
try:
|
||||
key = self.credentials.resolve(provider.credential_id)
|
||||
except CredentialStoreError as exc:
|
||||
raise ProviderError("PROVIDER_CREDENTIAL_UNAVAILABLE", "Provider credential is unavailable.") from exc
|
||||
if provider.credential_id and not key:
|
||||
raise ProviderError("PROVIDER_CREDENTIAL_MISSING", "Provider credential is not configured.")
|
||||
url = (provider.base_url or "https://api.openai.com/v1").rstrip("/") + binding.endpoint
|
||||
return url, {"Authorization": f"Bearer {key}"} if key else {}
|
||||
|
||||
async def _request(self, binding: ModelBinding, *, remote: tuple[str, dict[str, str]] | None = None, **kwargs) -> tuple[dict, str]:
|
||||
url, headers = remote or self._remote(binding)
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=30, transport=self.transport) as client:
|
||||
async with client.stream("POST", url, headers=headers, **kwargs) as response:
|
||||
response.raise_for_status()
|
||||
body = bytearray()
|
||||
async for chunk in response.aiter_bytes():
|
||||
body.extend(chunk)
|
||||
if len(body) > MAX_RESPONSE_BYTES:
|
||||
raise invalid_response()
|
||||
data = json.loads(body)
|
||||
except httpx.TimeoutException as exc:
|
||||
raise ProviderError("PROVIDER_TIMEOUT", "Model API timed out.") from exc
|
||||
except httpx.HTTPStatusError as exc:
|
||||
code = {401: "PROVIDER_AUTH_FAILED", 403: "PROVIDER_AUTH_FAILED", 404: "MODEL_NOT_FOUND", 429: "PROVIDER_RATE_LIMITED"}.get(exc.response.status_code, "PROVIDER_UNAVAILABLE")
|
||||
raise ProviderError(code, f"Model API returned HTTP {exc.response.status_code}.") from exc
|
||||
except (httpx.HTTPError, httpx.InvalidURL) as exc:
|
||||
raise ProviderError("PROVIDER_UNAVAILABLE", "Model API is unavailable.") from exc
|
||||
except (ValueError, UnicodeError) as exc:
|
||||
raise invalid_response() from exc
|
||||
if not isinstance(data, dict) or data.get("error"):
|
||||
raise invalid_response()
|
||||
return data, url
|
||||
|
||||
async def embed(self, texts: list[str]) -> EmbeddingResult:
|
||||
config = self.configuration()
|
||||
binding = config.embedding
|
||||
record_embedding(route_version=config.version,
|
||||
requested_route=binding.model_dump() if binding else None)
|
||||
reason = None
|
||||
if binding and texts:
|
||||
try:
|
||||
vectors = []
|
||||
dimension = binding.dimensions
|
||||
# Freeze the origin across batches, even if the user edits the provider.
|
||||
remote = self._remote(binding)
|
||||
for start in range(0, len(texts), 32):
|
||||
batch = texts[start:start + 32]
|
||||
payload = {"model": binding.model, "input": batch, "encoding_format": "float"}
|
||||
if binding.dimensions is not None:
|
||||
payload["dimensions"] = binding.dimensions
|
||||
data, url = await self._request(binding, remote=remote, json=payload)
|
||||
items = data.get("data")
|
||||
if not isinstance(items, list) or len(items) != len(batch):
|
||||
raise invalid_response()
|
||||
indexed = {}
|
||||
for item in items:
|
||||
if not isinstance(item, dict):
|
||||
raise invalid_response()
|
||||
index, vector = item.get("index"), item.get("embedding")
|
||||
if type(index) is not int or index in indexed or not 0 <= index < len(batch):
|
||||
raise invalid_response()
|
||||
if not isinstance(vector, list) or not 1 <= len(vector) <= 16384:
|
||||
raise invalid_response()
|
||||
if any(not finite_number(value) for value in vector):
|
||||
raise invalid_response()
|
||||
dimension = dimension or len(vector)
|
||||
norm = math.hypot(*vector)
|
||||
if len(vector) != dimension or not norm or not math.isfinite(norm):
|
||||
raise invalid_response()
|
||||
indexed[index] = [value / norm for value in vector]
|
||||
vectors.extend(indexed[index] for index in range(len(batch)))
|
||||
identity = json.dumps([url, binding.model, dimension], separators=(",", ":"))
|
||||
return EmbeddingResult(vectors=vectors, source="api", dimensions=dimension,
|
||||
model_id="api-" + hashlib.sha256(identity.encode()).hexdigest())
|
||||
except ProviderError as exc:
|
||||
reason = exc.code
|
||||
vectors = await self.local_embedding.embed_documents(texts)
|
||||
return EmbeddingResult(vectors=vectors, source="local", model_id=self.local_embedding.model_id,
|
||||
dimensions=self.local_embedding.dim, fallback_reason=reason)
|
||||
|
||||
@staticmethod
|
||||
def _media_file(path: Path):
|
||||
try:
|
||||
handle = path.open("rb")
|
||||
except OSError as exc:
|
||||
raise ApiError(404, "ATTACHMENT_NOT_FOUND", "Audio attachment was not found.") from exc
|
||||
import os
|
||||
if not 0 < os.fstat(handle.fileno()).st_size <= MAX_MEDIA_BYTES:
|
||||
handle.close()
|
||||
raise ApiError(413, "ATTACHMENT_TOO_LARGE", "Audio attachment must be between 1 byte and 25 MiB.")
|
||||
return handle
|
||||
|
||||
async def transcribe(self, source: Path, language: str | None) -> RoutedTranscript:
|
||||
binding = self.configuration().transcription
|
||||
if binding is None:
|
||||
with self._media_file(source):
|
||||
pass
|
||||
reason = None
|
||||
if binding:
|
||||
try:
|
||||
fields = {"model": binding.model}
|
||||
if language:
|
||||
fields["language"] = language
|
||||
with self._media_file(source) as handle:
|
||||
data, _ = await self._request(binding, data=fields,
|
||||
files={"file": (source.name, handle, "application/octet-stream")})
|
||||
text = data.get("text")
|
||||
if not isinstance(text, str) or not text.strip():
|
||||
raise invalid_response()
|
||||
return RoutedTranscript(text=text, source="api")
|
||||
except ProviderError as exc:
|
||||
reason = exc.code
|
||||
try:
|
||||
text = await self.local_speech.transcribe(source, language)
|
||||
if not isinstance(text, str) or not text.strip():
|
||||
raise ProviderError("LOCAL_MODEL_INVALID_RESPONSE", "Local transcription was empty.")
|
||||
return RoutedTranscript(text=text, source="local", fallback_reason=reason)
|
||||
except ProviderError as exc:
|
||||
raise ApiError(503, exc.code, exc.message, {"fallback_reason": reason}) from exc
|
||||
|
||||
async def match_speakers(self, source: Path, reference: Path) -> SpeakerMatchResult:
|
||||
binding = self.configuration().speaker_matching
|
||||
if binding is None:
|
||||
with self._media_file(source), self._media_file(reference):
|
||||
pass
|
||||
reason = None
|
||||
if binding:
|
||||
try:
|
||||
# Explicit application contract, not an OpenAI-standard endpoint.
|
||||
with self._media_file(source) as audio, self._media_file(reference) as sample:
|
||||
data, _ = await self._request(binding, data={"model": binding.model}, files={
|
||||
"file": (source.name, audio, "application/octet-stream"),
|
||||
"reference_file": (reference.name, sample, "application/octet-stream"),
|
||||
})
|
||||
score = data.get("score")
|
||||
if not finite_number(score) or not 0 <= score <= 1:
|
||||
raise invalid_response()
|
||||
return SpeakerMatchResult(score=score, source="api")
|
||||
except ProviderError as exc:
|
||||
reason = exc.code
|
||||
try:
|
||||
score = await self.local_speech.match(source, reference)
|
||||
if not finite_number(score) or not 0 <= score <= 1:
|
||||
raise ProviderError("LOCAL_MODEL_INVALID_RESPONSE", "Local speaker matching was invalid.")
|
||||
return SpeakerMatchResult(score=score, source="local", fallback_reason=reason)
|
||||
except ProviderError as exc:
|
||||
raise ApiError(503, exc.code, exc.message, {"fallback_reason": reason}) from exc
|
||||
@@ -0,0 +1,47 @@
|
||||
"""Keep internal namespaced tools compatible with providers' 64-character names."""
|
||||
import hashlib
|
||||
import re
|
||||
from functools import wraps
|
||||
|
||||
from app.contracts import MessageRole, ModelRequest
|
||||
|
||||
|
||||
def prepare_tool_names(request: ModelRequest) -> tuple[ModelRequest, dict[str, str]]:
|
||||
names = {tool.name for tool in request.tools}
|
||||
for message in request.messages:
|
||||
names.update(call.name for call in message.tool_calls)
|
||||
if message.role == MessageRole.tool and message.name:
|
||||
names.add(message.name)
|
||||
mapping = {name: name for name in names if re.fullmatch(r"[A-Za-z0-9_-]{1,64}", name)}
|
||||
used = set(mapping)
|
||||
for name in sorted(names - mapping.keys()):
|
||||
salt = 0
|
||||
while True:
|
||||
alias = "tool_" + hashlib.sha256(f"{name}:{salt}".encode()).hexdigest()[:56]
|
||||
if alias not in used:
|
||||
break
|
||||
salt += 1
|
||||
mapping[name] = alias
|
||||
used.add(alias)
|
||||
if all(name == alias for name, alias in mapping.items()):
|
||||
return request, {}
|
||||
wire = request.model_copy(deep=True)
|
||||
for tool in wire.tools:
|
||||
tool.name = mapping[tool.name]
|
||||
for message in wire.messages:
|
||||
for call in message.tool_calls:
|
||||
call.name = mapping[call.name]
|
||||
if message.role == MessageRole.tool and message.name:
|
||||
message.name = mapping[message.name]
|
||||
return wire, {alias: name for name, alias in mapping.items()}
|
||||
|
||||
|
||||
def mapped_tool_names(complete):
|
||||
@wraps(complete)
|
||||
async def wrapped(self, request: ModelRequest):
|
||||
wire, originals = prepare_tool_names(request)
|
||||
turn = await complete(self, wire)
|
||||
for call in turn.tool_calls:
|
||||
call.name = originals.get(call.name, call.name)
|
||||
return turn
|
||||
return wrapped
|
||||
@@ -273,11 +273,15 @@ def update_note_location(
|
||||
raise LookupError(note_id)
|
||||
|
||||
|
||||
def fts_search_page(
|
||||
*,
|
||||
_FTS_FROM = """
|
||||
FROM blocks_fts
|
||||
JOIN blocks AS b ON b.block_id = blocks_fts.block_id
|
||||
JOIN notes AS n ON n.note_id = b.note_id
|
||||
"""
|
||||
|
||||
|
||||
def _fts_where(
|
||||
match: str,
|
||||
limit: int,
|
||||
offset: int,
|
||||
folders: list[str],
|
||||
note_ids: list[str],
|
||||
tags: list[str],
|
||||
@@ -285,8 +289,11 @@ def fts_search_page(
|
||||
created_to: datetime | None,
|
||||
updated_from: datetime | None,
|
||||
updated_to: datetime | None,
|
||||
) -> tuple[list[FtsHit], int]:
|
||||
"""执行带元数据过滤的 FTS 精确分页,并返回过滤后的完整命中数。"""
|
||||
) -> tuple[str, list[object]]:
|
||||
"""构建 FTS 过滤 WHERE 子句(不含 WHERE 关键字),返回 (where_sql, params)。
|
||||
|
||||
fts_search_page 与 fts_score_bounds 共用,保证计数与取数口径一致。
|
||||
"""
|
||||
where = ["blocks_fts MATCH ?"]
|
||||
params: list[object] = [match]
|
||||
|
||||
@@ -317,22 +324,44 @@ def fts_search_page(
|
||||
where.append(f"julianday({column}) <= julianday(?)")
|
||||
params.append(_iso(upper))
|
||||
|
||||
from_sql = """
|
||||
FROM blocks_fts
|
||||
JOIN blocks AS b ON b.block_id = blocks_fts.block_id
|
||||
JOIN notes AS n ON n.note_id = b.note_id
|
||||
return " AND ".join(where), params
|
||||
|
||||
|
||||
def fts_search_page(
|
||||
*,
|
||||
match: str,
|
||||
limit: int,
|
||||
offset: int,
|
||||
folders: list[str],
|
||||
note_ids: list[str],
|
||||
tags: list[str],
|
||||
created_from: datetime | None,
|
||||
created_to: datetime | None,
|
||||
updated_from: datetime | None,
|
||||
updated_to: datetime | None,
|
||||
bm25_max: float | None = None,
|
||||
) -> tuple[list[FtsHit], int]:
|
||||
"""执行带元数据过滤的 FTS 精确分页,并返回过滤后的完整命中数。
|
||||
|
||||
bm25_max 非空时按 bm25 截止值过滤(用于阈值过滤的精确分页),计数与取数同口径。
|
||||
"""
|
||||
where_sql = " AND ".join(where)
|
||||
where_sql, params = _fts_where(
|
||||
match, folders, note_ids, tags,
|
||||
created_from, created_to, updated_from, updated_to,
|
||||
)
|
||||
if bm25_max is not None:
|
||||
where_sql += " AND bm25(blocks_fts) <= ?"
|
||||
params.append(bm25_max)
|
||||
|
||||
conn = connect()
|
||||
try:
|
||||
total = conn.execute(
|
||||
f"SELECT COUNT(*) {from_sql} WHERE {where_sql}", params
|
||||
f"SELECT COUNT(*) {_FTS_FROM} WHERE {where_sql}", params
|
||||
).fetchone()[0]
|
||||
rows = conn.execute(
|
||||
f"""
|
||||
SELECT blocks_fts.block_id, blocks_fts.note_id, bm25(blocks_fts) AS rank
|
||||
{from_sql}
|
||||
{_FTS_FROM}
|
||||
WHERE {where_sql}
|
||||
ORDER BY rank
|
||||
LIMIT ? OFFSET ?
|
||||
@@ -348,6 +377,45 @@ def fts_search_page(
|
||||
conn.close()
|
||||
|
||||
|
||||
def fts_score_bounds(
|
||||
*,
|
||||
match: str,
|
||||
folders: list[str],
|
||||
note_ids: list[str],
|
||||
tags: list[str],
|
||||
created_from: datetime | None,
|
||||
created_to: datetime | None,
|
||||
updated_from: datetime | None,
|
||||
updated_to: datetime | None,
|
||||
) -> tuple[float, float] | None:
|
||||
"""返回 metadata 过滤后的 FTS 命中集里 bm25 的 (min, max),无命中时返回 None。
|
||||
|
||||
用于阈值过滤:min-max 归一化是 bm25 的线性函数,据此可把阈值换算为 bm25 截止值。
|
||||
"""
|
||||
where_sql, params = _fts_where(
|
||||
match, folders, note_ids, tags,
|
||||
created_from, created_to, updated_from, updated_to,
|
||||
)
|
||||
conn = connect()
|
||||
try:
|
||||
# bm25() 不能作为聚合函数参数,也不能用在被聚合的子查询里;改用 ORDER BY 取首尾两行
|
||||
lo_row = conn.execute(
|
||||
f"SELECT bm25(blocks_fts) AS rank {_FTS_FROM} WHERE {where_sql}"
|
||||
" ORDER BY rank ASC LIMIT 1",
|
||||
params,
|
||||
).fetchone()
|
||||
if lo_row is None or lo_row["rank"] is None:
|
||||
return None
|
||||
hi_row = conn.execute(
|
||||
f"SELECT bm25(blocks_fts) AS rank {_FTS_FROM} WHERE {where_sql}"
|
||||
" ORDER BY rank DESC LIMIT 1",
|
||||
params,
|
||||
).fetchone()
|
||||
return (float(lo_row["rank"]), float(hi_row["rank"]))
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def get_block_hits(block_ids: list[str]) -> list[BlockHit]:
|
||||
if not block_ids:
|
||||
return []
|
||||
@@ -392,15 +460,17 @@ def get_index_meta() -> dict[str, str]:
|
||||
conn.close()
|
||||
|
||||
|
||||
def clear_all() -> None:
|
||||
"""清空元数据、Block 与 FTS5(重建索引用,向量由 VectorStore.clear 处理)。"""
|
||||
conn = connect()
|
||||
def clear_all(*, conn: sqlite3.Connection | None = None) -> None:
|
||||
"""Clear rebuildable metadata using the caller's transaction when provided."""
|
||||
owns = conn is None
|
||||
conn = conn or connect()
|
||||
try:
|
||||
with transaction(conn):
|
||||
with transaction(conn) if owns else nullcontext():
|
||||
conn.execute("DELETE FROM blocks_fts")
|
||||
conn.execute("DELETE FROM blocks")
|
||||
conn.execute("DELETE FROM notes")
|
||||
finally:
|
||||
if owns:
|
||||
conn.close()
|
||||
|
||||
|
||||
|
||||
@@ -19,6 +19,7 @@ class EmbeddingProvider(Protocol):
|
||||
"""统一 Embedding 接口(与文档一致)。"""
|
||||
|
||||
model_id: str
|
||||
version: str
|
||||
dim: int
|
||||
|
||||
async def embed_documents(self, texts: list[str]) -> list[list[float]]: ...
|
||||
@@ -33,6 +34,7 @@ class HashEmbeddingProvider:
|
||||
"""
|
||||
|
||||
model_id = "hash-v1"
|
||||
version = "1"
|
||||
dim = EMBEDDING_DIM
|
||||
|
||||
async def embed_documents(self, texts: list[str]) -> list[list[float]]:
|
||||
|
||||
@@ -22,6 +22,8 @@ from app.repository import BlockHit
|
||||
from app.retrieval.embedding import EmbeddingProvider, HashEmbeddingProvider
|
||||
from app.retrieval.hybrid import normalize_scores, rrf_fuse
|
||||
from app.retrieval.reranker import LexicalReranker, RankedCandidate, RerankerProvider
|
||||
from app.retrieval import routed_vectors
|
||||
from app.retrieval.provenance import record_embedding
|
||||
from app.retrieval.vectorstore import SqliteVecStore, VectorStore
|
||||
from app.textutils import make_snippet, match_query
|
||||
|
||||
@@ -39,10 +41,15 @@ class RetrievalEngine:
|
||||
embedding: EmbeddingProvider,
|
||||
reranker: RerankerProvider,
|
||||
vector_store: VectorStore,
|
||||
*,
|
||||
route_embeddings: bool = False,
|
||||
) -> None:
|
||||
self.embedding = embedding
|
||||
self.reranker = reranker
|
||||
self.vector_store = vector_store
|
||||
# Only the production instance opts in. Replaced test dependencies must
|
||||
# remain authoritative, including monkeypatches on the singleton.
|
||||
self._routed_defaults = (embedding, vector_store) if route_embeddings else None
|
||||
|
||||
async def search(self, request: SearchRequest) -> SearchResponse:
|
||||
if request.mode == SearchMode.fts:
|
||||
@@ -56,7 +63,7 @@ class RetrievalEngine:
|
||||
# 候选池至少覆盖本次请求的 offset+limit,保证分页能取到目标页;设上限防内存失控
|
||||
window = min(request.offset + request.limit, MAX_CANDIDATE_POOL)
|
||||
pool_size = max(CANDIDATE_POOL, window)
|
||||
# 带过滤时放大召回;FTS 则一次性取全量命中(≤FTS_FETCH_LIMIT)避免截断漏召回
|
||||
# 带过滤时放大召回,缓解「先截断候选池再过滤」造成的漏召回
|
||||
recall = min(pool_size * OVERSCAN_FACTOR, MAX_CANDIDATE_POOL) if has_filters else pool_size
|
||||
|
||||
# 1. 按模式收集候选(FTS 与 Vector 各产出「按相关性降序」的 block_id 列表)
|
||||
@@ -74,8 +81,19 @@ class RetrievalEngine:
|
||||
fts_scores = {h.block_id: -h.bm25 for h in fts_hits}
|
||||
|
||||
if request.mode in (SearchMode.vector, SearchMode.hybrid):
|
||||
record_embedding(source="unavailable")
|
||||
vec_hits = None
|
||||
if (
|
||||
self._routed_defaults is not None
|
||||
and self.embedding is self._routed_defaults[0]
|
||||
and self.vector_store is self._routed_defaults[1]
|
||||
):
|
||||
vec_hits = await routed_vectors.search_remote(request.query, top_k=recall)
|
||||
if vec_hits is None:
|
||||
query_vec = await self.embedding.embed_query(request.query)
|
||||
vec_hits = await self.vector_store.search(query_vec, top_k=recall)
|
||||
record_embedding(source="local", model_id=self.embedding.model_id,
|
||||
dimensions=self.embedding.dim, version=self.embedding.version)
|
||||
vec_ranked = [v.id for v in vec_hits]
|
||||
vec_scores = {v.id: v.score for v in vec_hits}
|
||||
|
||||
@@ -84,7 +102,7 @@ class RetrievalEngine:
|
||||
elif request.mode == SearchMode.vector:
|
||||
candidate_scores = vec_scores
|
||||
else: # hybrid:RRF 融合
|
||||
candidate_scores = rrf_fuse([fts_ranked, vec_ranked])
|
||||
candidate_scores = rrf_fuse([fts_ranked, vec_ranked], k=request.rrf_k)
|
||||
|
||||
if not candidate_scores:
|
||||
return self._empty(request)
|
||||
@@ -97,14 +115,23 @@ class RetrievalEngine:
|
||||
if not filtered:
|
||||
return self._empty(request)
|
||||
|
||||
# 4. 排序 / 精排
|
||||
# 4. 排序 / 精排:hybrid 先按融合分预排序,再对前 rerank_candidates 个候选做精排,
|
||||
# 剩余候选按融合分排在精排结果之后;rerank=False 时跳过精排直接按融合分排序。
|
||||
if request.mode == SearchMode.hybrid:
|
||||
pre_sorted = sorted(filtered, key=lambda h: -candidate_scores[h.block_id])
|
||||
if request.rerank:
|
||||
limit = request.rerank_candidates
|
||||
pool = pre_sorted if limit is None else pre_sorted[:limit]
|
||||
rest = [] if limit is None else pre_sorted[limit:]
|
||||
candidates = [
|
||||
RankedCandidate(block_id=h.block_id, score=candidate_scores[h.block_id], text=h.content)
|
||||
for h in filtered
|
||||
for h in pool
|
||||
]
|
||||
ranked = await self.reranker.rerank(request.query, candidates)
|
||||
ordered = [(c.block_id, c.score) for c in ranked]
|
||||
ordered += [(h.block_id, candidate_scores[h.block_id]) for h in rest]
|
||||
else:
|
||||
ordered = [(h.block_id, candidate_scores[h.block_id]) for h in pre_sorted]
|
||||
else:
|
||||
ordered = sorted(
|
||||
((h.block_id, candidate_scores[h.block_id]) for h in filtered),
|
||||
@@ -112,8 +139,10 @@ class RetrievalEngine:
|
||||
)
|
||||
|
||||
ordered = normalize_scores(ordered)
|
||||
# score_threshold:归一化后过滤低分结果(默认 0 不过滤)
|
||||
ordered = [(bid, score) for bid, score in ordered if score >= request.score_threshold]
|
||||
|
||||
# 5. 分页:total = 过滤后候选集大小。fts 已取全量(≤FTS_FETCH_LIMIT)故为真实命中数;
|
||||
# 5. 分页:total = 过滤后候选集大小。fts 走数据库精确分页,total 为真实命中数;
|
||||
# vector/hybrid 为 KNN 候选集,无全局 total。
|
||||
total = len(ordered)
|
||||
page = ordered[request.offset : request.offset + request.limit]
|
||||
@@ -126,11 +155,41 @@ class RetrievalEngine:
|
||||
)
|
||||
|
||||
def _search_fts(self, request: SearchRequest) -> SearchResponse:
|
||||
"""FTS 专用路径:过滤、COUNT 与分页全部在 SQLite 中完成。"""
|
||||
"""FTS 专用路径:在数据库侧完成过滤、计数与分页,不取全量后再截断。
|
||||
|
||||
阈值过滤时,min-max 归一化是 bm25 的线性函数,据此把 score_threshold 换算为
|
||||
bm25 截止值(bm25_max),使过滤、计数与分页口径一致;无阈值时走数据库原生分页,
|
||||
total 始终为过滤后的真实命中数,不再受固定截断影响。
|
||||
"""
|
||||
match = match_query(request.query)
|
||||
if not match:
|
||||
return self._empty(request)
|
||||
|
||||
bounds = repository.fts_score_bounds(
|
||||
match=match,
|
||||
folders=request.folders,
|
||||
note_ids=request.note_ids,
|
||||
tags=request.tags,
|
||||
created_from=request.created_from,
|
||||
created_to=request.created_to,
|
||||
updated_from=request.updated_from,
|
||||
updated_to=request.updated_to,
|
||||
)
|
||||
if bounds is None:
|
||||
return self._empty(request)
|
||||
|
||||
lo, hi = bounds
|
||||
span = hi - lo
|
||||
bm25_max: float | None = None
|
||||
if request.score_threshold > 0:
|
||||
if span == 0:
|
||||
# 全部命中 bm25 相同,归一化后皆为 1.0;阈值超过 1.0 时无命中
|
||||
if request.score_threshold > 1.0:
|
||||
return self._empty(request)
|
||||
else:
|
||||
# norm = (hi - bm25) / span;norm >= threshold ⟺ bm25 <= hi - threshold * span
|
||||
bm25_max = hi - request.score_threshold * span
|
||||
|
||||
fts_hits, total = repository.fts_search_page(
|
||||
match=match,
|
||||
limit=request.limit,
|
||||
@@ -142,19 +201,29 @@ class RetrievalEngine:
|
||||
created_to=request.created_to,
|
||||
updated_from=request.updated_from,
|
||||
updated_to=request.updated_to,
|
||||
bm25_max=bm25_max,
|
||||
)
|
||||
if not fts_hits:
|
||||
# 本页无结果:offset 越过末页时 total 仍为真实命中数(>0),需保留而非归零
|
||||
return SearchResponse(
|
||||
query=request.query,
|
||||
mode=request.mode,
|
||||
items=[],
|
||||
page=PageMeta(total=total, limit=request.limit, offset=request.offset),
|
||||
)
|
||||
|
||||
hits = {h.block_id: h for h in repository.get_block_hits([hit.block_id for hit in fts_hits])}
|
||||
ordered = normalize_scores(
|
||||
[(hit.block_id, -hit.bm25) for hit in fts_hits if hit.block_id in hits]
|
||||
)
|
||||
items = [self._build_result(hits[block_id], request, score) for block_id, score in ordered]
|
||||
# 分数按全局 bm25 上下界归一化(与取全量后 normalize_scores 等价),保证跨页一致
|
||||
span = hi - lo
|
||||
if span == 0:
|
||||
ordered = [(hit.block_id, 1.0) for hit in fts_hits]
|
||||
else:
|
||||
ordered = [(hit.block_id, round((hi - hit.bm25) / span, 6)) for hit in fts_hits]
|
||||
hits = {h.block_id: h for h in repository.get_block_hits([bid for bid, _ in ordered])}
|
||||
items = [
|
||||
self._build_result(hits[block_id], request, score)
|
||||
for block_id, score in ordered
|
||||
if block_id in hits
|
||||
]
|
||||
return SearchResponse(
|
||||
query=request.query,
|
||||
mode=request.mode,
|
||||
@@ -217,4 +286,6 @@ def _utc(dt: datetime) -> datetime:
|
||||
|
||||
|
||||
# 默认引擎实例:轻量实现跑通链路,后续可替换真实模型实现
|
||||
engine = RetrievalEngine(HashEmbeddingProvider(), LexicalReranker(), SqliteVecStore())
|
||||
engine = RetrievalEngine(
|
||||
HashEmbeddingProvider(), LexicalReranker(), SqliteVecStore(), route_embeddings=True,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
"""Task-local observations of the embedding path actually used by a search."""
|
||||
from contextlib import contextmanager
|
||||
from contextvars import ContextVar
|
||||
|
||||
_observation: ContextVar[dict | None] = ContextVar("embedding_observation", default=None)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def capture_embedding():
|
||||
result = {"source": "not_used"}
|
||||
token = _observation.set(result)
|
||||
try:
|
||||
yield result
|
||||
finally:
|
||||
_observation.reset(token)
|
||||
|
||||
|
||||
def record_embedding(**fields) -> None:
|
||||
result = _observation.get()
|
||||
if result is not None:
|
||||
result.update(fields)
|
||||
@@ -24,6 +24,7 @@ class RerankerProvider(Protocol):
|
||||
"""统一 Reranker 接口:输入候选块,输出按相关性重排后的候选块。"""
|
||||
|
||||
model_id: str
|
||||
version: str
|
||||
|
||||
async def rerank(self, query: str, candidates: list[RankedCandidate]) -> list[RankedCandidate]: ...
|
||||
|
||||
@@ -32,6 +33,7 @@ class LexicalReranker:
|
||||
"""轻量精排:query 与块正文的词面重叠度,与归一化后的原始分数加权求和。"""
|
||||
|
||||
model_id = "lexical-v1"
|
||||
version = "1"
|
||||
|
||||
def __init__(self, lexical_weight: float = 0.5) -> None:
|
||||
self.lexical_weight = lexical_weight
|
||||
|
||||
@@ -0,0 +1,201 @@
|
||||
"""Optional API embeddings, isolated from the stable hash/sqlite-vec index.
|
||||
|
||||
The runtime's model_id is the authoritative space ID (including provider URL,
|
||||
endpoint, model and dimensions); equal dimensions alone never imply compatibility.
|
||||
This phase uses a lazy, rebuildable SQLite side table instead of a schema migration.
|
||||
Search scans only current blocks in one database snapshot and requires complete
|
||||
coverage. Cosine ranking costs O(blocks * dimensions) with an O(top_k) heap; this
|
||||
small-vault implementation should become a per-space ANN index at larger scale.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import heapq
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import sqlite3
|
||||
from dataclasses import dataclass
|
||||
from typing import Protocol
|
||||
|
||||
from app.database.db import connect, transaction
|
||||
from app.retrieval.vectorstore import VectorHit
|
||||
from app.retrieval.provenance import record_embedding
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class EmbeddingResult(Protocol):
|
||||
vectors: list[list[float]]
|
||||
source: str
|
||||
model_id: str
|
||||
dimensions: int
|
||||
fallback_reason: str | None
|
||||
|
||||
|
||||
class EmbeddingRuntime(Protocol):
|
||||
async def embed(self, texts: list[str]) -> EmbeddingResult: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RemoteEmbeddings:
|
||||
space_id: str
|
||||
dimensions: int
|
||||
vectors: list[list[float]]
|
||||
|
||||
|
||||
def get_model_routing() -> EmbeddingRuntime | None:
|
||||
"""Lazy integration hook; tests can inject a runtime without any network I/O."""
|
||||
from app.container import container
|
||||
|
||||
return getattr(container, "model_routing", None)
|
||||
|
||||
|
||||
def _unit_vector(vector: list[float], dimensions: int) -> list[float]:
|
||||
if len(vector) != dimensions:
|
||||
raise ValueError("embedding dimension mismatch")
|
||||
if any(isinstance(value, bool) or not isinstance(value, (int, float)) for value in vector):
|
||||
raise ValueError("embedding must be numeric")
|
||||
if not all(math.isfinite(value) for value in vector):
|
||||
raise ValueError("embedding must be finite")
|
||||
scale = max(abs(value) for value in vector)
|
||||
if scale == 0:
|
||||
raise ValueError("embedding must be nonzero")
|
||||
# Scaling first avoids overflow/underflow for finite but extreme API values.
|
||||
scaled = [value / scale for value in vector]
|
||||
norm = math.sqrt(math.fsum(value * value for value in scaled))
|
||||
return [value / norm for value in scaled]
|
||||
|
||||
|
||||
async def embed_remote(texts: list[str]) -> RemoteEmbeddings | None:
|
||||
"""Return validated API vectors, or None to use the caller's local baseline.
|
||||
|
||||
Do not use the runtime's local result: the caller may have injected its own
|
||||
embedding/store pair. Exception deliberately excludes cancellation.
|
||||
"""
|
||||
if not texts:
|
||||
return None
|
||||
try:
|
||||
runtime = get_model_routing()
|
||||
if runtime is None:
|
||||
return None
|
||||
result = await runtime.embed(texts)
|
||||
if result.source != "api":
|
||||
record_embedding(fallback_reason=result.fallback_reason)
|
||||
return None
|
||||
if not isinstance(result.model_id, str) or not result.model_id or result.model_id == "hash-v1":
|
||||
raise ValueError("API embedding needs a distinct space ID")
|
||||
if type(result.dimensions) is not int or result.dimensions <= 0:
|
||||
raise ValueError("invalid embedding dimensions")
|
||||
if len(result.vectors) != len(texts):
|
||||
raise ValueError("embedding count mismatch")
|
||||
return RemoteEmbeddings(
|
||||
space_id=result.model_id,
|
||||
dimensions=result.dimensions,
|
||||
vectors=[_unit_vector(vector, result.dimensions) for vector in result.vectors],
|
||||
)
|
||||
except Exception as exc:
|
||||
# Avoid logging provider exceptions containing credentials or note text.
|
||||
record_embedding(fallback_reason="REMOTE_EMBEDDING_UNAVAILABLE")
|
||||
logger.warning("Remote embedding unavailable (%s); using local index", type(exc).__name__)
|
||||
return None
|
||||
|
||||
|
||||
def _ensure_table(conn: sqlite3.Connection) -> None:
|
||||
conn.execute("""
|
||||
CREATE TABLE IF NOT EXISTS routed_block_vectors (
|
||||
space_id TEXT NOT NULL,
|
||||
block_id TEXT NOT NULL REFERENCES blocks(block_id) ON DELETE CASCADE,
|
||||
dimensions INTEGER NOT NULL CHECK (dimensions > 0),
|
||||
vector TEXT NOT NULL,
|
||||
PRIMARY KEY (space_id, block_id)
|
||||
)
|
||||
""")
|
||||
conn.execute("""
|
||||
CREATE INDEX IF NOT EXISTS routed_block_vectors_block_id
|
||||
ON routed_block_vectors(block_id)
|
||||
""")
|
||||
|
||||
|
||||
def store_remote(
|
||||
conn: sqlite3.Connection, block_ids: list[str], batch: RemoteEmbeddings | None,
|
||||
) -> None:
|
||||
"""Best-effort side-index write inside the caller's metadata transaction.
|
||||
|
||||
A savepoint prevents partial remote batches and isolates storage failures from
|
||||
note saving. Replacing/deleting blocks cascades all old spaces automatically.
|
||||
"""
|
||||
if batch is None:
|
||||
return
|
||||
try:
|
||||
conn.execute("SAVEPOINT routed_vectors_write")
|
||||
try:
|
||||
if len(block_ids) != len(batch.vectors):
|
||||
raise ValueError("block/vector count mismatch")
|
||||
_ensure_table(conn)
|
||||
conn.executemany(
|
||||
"""INSERT INTO routed_block_vectors (space_id, block_id, dimensions, vector)
|
||||
VALUES (?, ?, ?, ?)
|
||||
ON CONFLICT (space_id, block_id) DO UPDATE SET
|
||||
dimensions = excluded.dimensions, vector = excluded.vector""",
|
||||
[
|
||||
(batch.space_id, block_id, batch.dimensions, json.dumps(vector, allow_nan=False))
|
||||
for block_id, vector in zip(block_ids, batch.vectors)
|
||||
],
|
||||
)
|
||||
except BaseException:
|
||||
conn.execute("ROLLBACK TO routed_vectors_write")
|
||||
raise
|
||||
finally:
|
||||
conn.execute("RELEASE routed_vectors_write")
|
||||
except Exception as exc:
|
||||
logger.warning("Remote vector storage unavailable (%s); local index retained", type(exc).__name__)
|
||||
|
||||
|
||||
async def search_remote(query: str, *, top_k: int) -> list[VectorHit] | None:
|
||||
"""None means fallback, including any missing/invalid current-block vector.
|
||||
|
||||
Read coverage and vectors together so concurrent note updates cannot produce
|
||||
an apparently complete subset. Never fill missing remote hits with local hits.
|
||||
"""
|
||||
batch = await embed_remote([query])
|
||||
if batch is None:
|
||||
return None
|
||||
record_embedding(attempted_space={"model_id": batch.space_id, "dimensions": batch.dimensions})
|
||||
try:
|
||||
conn = connect()
|
||||
try:
|
||||
with transaction(conn):
|
||||
exists = conn.execute(
|
||||
"SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'routed_block_vectors'"
|
||||
).fetchone()
|
||||
if exists is None:
|
||||
record_embedding(fallback_reason="REMOTE_INDEX_MISSING")
|
||||
return None
|
||||
rows = conn.execute(
|
||||
"""SELECT b.block_id, r.vector
|
||||
FROM blocks AS b
|
||||
LEFT JOIN routed_block_vectors AS r
|
||||
ON r.block_id = b.block_id AND r.space_id = ? AND r.dimensions = ?
|
||||
ORDER BY b.block_id""",
|
||||
(batch.space_id, batch.dimensions),
|
||||
)
|
||||
|
||||
def hits():
|
||||
for row in rows:
|
||||
if row["vector"] is None:
|
||||
raise ValueError("remote space has incomplete block coverage")
|
||||
vector = _unit_vector(json.loads(row["vector"]), batch.dimensions)
|
||||
score = math.fsum(a * b for a, b in zip(batch.vectors[0], vector))
|
||||
yield VectorHit(id=row["block_id"], score=max(0.0, min(1.0, score)))
|
||||
|
||||
result = heapq.nlargest(top_k, hits(), key=lambda hit: hit.score)
|
||||
record_embedding(source="api", model_id=batch.space_id,
|
||||
dimensions=batch.dimensions, fallback_reason=None)
|
||||
return result
|
||||
finally:
|
||||
conn.close()
|
||||
except Exception as exc:
|
||||
record_embedding(fallback_reason="REMOTE_INDEX_UNAVAILABLE")
|
||||
logger.debug("Remote vector search unavailable (%s); using local index", type(exc).__name__)
|
||||
return None
|
||||
@@ -35,6 +35,7 @@ class VectorStore(Protocol):
|
||||
async def upsert(self, records: list[VectorRecord]) -> None: ...
|
||||
async def delete(self, ids: list[str]) -> None: ...
|
||||
async def search(self, vector: list[float], *, top_k: int) -> list[VectorHit]: ...
|
||||
async def count(self) -> int: ...
|
||||
|
||||
|
||||
class SqliteVecStore:
|
||||
@@ -85,10 +86,19 @@ class SqliteVecStore:
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
async def clear(self) -> None:
|
||||
conn = connect()
|
||||
async def clear(self, *, conn: sqlite3.Connection | None = None) -> None:
|
||||
owns = conn is None
|
||||
conn = conn or connect()
|
||||
try:
|
||||
with transaction(conn):
|
||||
with transaction(conn) if owns else nullcontext():
|
||||
conn.execute("DELETE FROM vec_blocks")
|
||||
finally:
|
||||
if owns:
|
||||
conn.close()
|
||||
|
||||
async def count(self) -> int:
|
||||
conn = connect()
|
||||
try:
|
||||
return conn.execute("SELECT COUNT(*) FROM vec_blocks").fetchone()[0]
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
@@ -1,17 +1,28 @@
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import aclosing
|
||||
from datetime import datetime, timezone
|
||||
from uuid import uuid4
|
||||
|
||||
from fastapi import APIRouter, Header, Query
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
from app.agent import AgentCapacityError, AgentRunNotFoundError
|
||||
from app.container import container
|
||||
from app.contracts import (
|
||||
AgentRun,
|
||||
AgentRunCreateRequest,
|
||||
AgentRunListResponse,
|
||||
AgentTraceResponse,
|
||||
ChatRequest,
|
||||
BenchmarkDatasetListResponse,
|
||||
BenchmarkEventType,
|
||||
BenchmarkKind,
|
||||
BenchmarkReport,
|
||||
BenchmarkRun,
|
||||
BenchmarkRunListResponse,
|
||||
BenchmarkStatus,
|
||||
RAGRunRequest,
|
||||
CredentialStatus,
|
||||
CredentialWriteRequest,
|
||||
ExtensionInstallRequest,
|
||||
@@ -21,8 +32,22 @@ from app.contracts import (
|
||||
IndexJob,
|
||||
IndexRebuildRequest,
|
||||
IndexStatus,
|
||||
McpServer,
|
||||
McpServerCreateRequest,
|
||||
McpServerListResponse,
|
||||
McpServerSecretStatus,
|
||||
McpServerSecretWriteRequest,
|
||||
McpServerTrustRequest,
|
||||
McpServerUpdateRequest,
|
||||
McpToolSummaryListResponse,
|
||||
ModelEvent,
|
||||
ModelEventType,
|
||||
EmbeddingRequest,
|
||||
EmbeddingResult,
|
||||
ModelRoutingConfig,
|
||||
ModelRoutingResponse,
|
||||
SpeakerMatchRequest,
|
||||
SpeakerMatchResult,
|
||||
Note,
|
||||
NoteCreateRequest,
|
||||
NoteListResponse,
|
||||
@@ -69,16 +94,19 @@ from app.contracts import (
|
||||
WorkspaceSnapshot,
|
||||
)
|
||||
from app.agent import AgentCapacityError, AgentRunNotFoundError
|
||||
from app.benchmarks import datasets as benchmark_datasets
|
||||
from app.benchmarks import service as benchmark_service
|
||||
from app.container import container
|
||||
from app.errors import ApiError
|
||||
from app.extensions import ExtensionError
|
||||
from app.providers.registry import ProviderNotFoundError
|
||||
from app.providers.factory import UnsupportedProviderError
|
||||
from app.extensions.mcp_registry import McpRegistryError
|
||||
from app.providers.base import ProviderError
|
||||
from app.providers.credentials import (
|
||||
CredentialStoreError,
|
||||
validate_provider_credential_id,
|
||||
)
|
||||
from app.providers.factory import UnsupportedProviderError
|
||||
from app.providers.registry import ProviderNotFoundError
|
||||
from app.retrieval.engine import engine
|
||||
from app.services import (
|
||||
index_service,
|
||||
@@ -87,10 +115,19 @@ from app.services import (
|
||||
transcription_service,
|
||||
workspace_service,
|
||||
)
|
||||
from app.services.attachment_service import attachment_path
|
||||
|
||||
router = APIRouter(prefix="/api")
|
||||
|
||||
|
||||
async def mcp_call_async(operation):
|
||||
"""Even registry reads can wait on lifecycle locks; keep all MCP work off the event loop."""
|
||||
try:
|
||||
return await asyncio.to_thread(operation)
|
||||
except McpRegistryError as exc:
|
||||
raise ApiError(exc.status_code, exc.code, exc.message) from exc
|
||||
|
||||
|
||||
def utc_now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
@@ -202,14 +239,21 @@ async def list_notes(
|
||||
folder: str | None = None,
|
||||
tag: str | None = None,
|
||||
) -> NoteListResponse:
|
||||
items, total = note_service.list_notes(limit=limit, offset=offset, folder=folder, tag=tag)
|
||||
return NoteListResponse(items=items, page=PageMeta(total=total, limit=limit, offset=offset))
|
||||
items, total = note_service.list_notes(
|
||||
limit=limit, offset=offset, folder=folder, tag=tag
|
||||
)
|
||||
return NoteListResponse(
|
||||
items=items, page=PageMeta(total=total, limit=limit, offset=offset)
|
||||
)
|
||||
|
||||
|
||||
@router.post("/notes", response_model=Note, tags=["Notes"])
|
||||
async def create_note(request: NoteCreateRequest) -> Note:
|
||||
return await note_service.create_note(
|
||||
title=request.title, markdown=request.markdown, folder=request.folder, tags=request.tags
|
||||
title=request.title,
|
||||
markdown=request.markdown,
|
||||
folder=request.folder,
|
||||
tags=request.tags,
|
||||
)
|
||||
|
||||
|
||||
@@ -217,7 +261,9 @@ async def create_note(request: NoteCreateRequest) -> Note:
|
||||
async def get_note(note_id: str) -> Note:
|
||||
note = await note_service.get_note(note_id)
|
||||
if note is None:
|
||||
raise ApiError(404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id})
|
||||
raise ApiError(
|
||||
404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id}
|
||||
)
|
||||
return note
|
||||
|
||||
|
||||
@@ -231,7 +277,9 @@ async def update_note(note_id: str, request: NoteUpdateRequest) -> Note:
|
||||
@router.delete("/notes/{note_id}", response_model=OperationResponse, tags=["Notes"])
|
||||
async def delete_note(note_id: str) -> OperationResponse:
|
||||
if not await note_service.delete_note(note_id):
|
||||
raise ApiError(404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id})
|
||||
raise ApiError(
|
||||
404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id}
|
||||
)
|
||||
return OperationResponse(status="completed", resource_id=note_id, message="deleted")
|
||||
|
||||
|
||||
@@ -266,16 +314,23 @@ async def chat(request: ChatRequest) -> StreamingResponse:
|
||||
provider = provider_or_404(request.provider_id)
|
||||
|
||||
async def stream() -> AsyncIterator[str]:
|
||||
sequence = 0
|
||||
try:
|
||||
async for event in provider.adapter.stream(request):
|
||||
async with aclosing(provider.adapter.stream(request)) as events:
|
||||
async for event in events:
|
||||
sequence = event.sequence + 1
|
||||
yield as_sse(event.event.value, event.model_dump_json())
|
||||
except Exception as exc:
|
||||
except Exception:
|
||||
error = ModelEvent(
|
||||
event=ModelEventType.error,
|
||||
data={"code": "PROVIDER_ERROR", "message": str(exc)},
|
||||
sequence=sequence,
|
||||
data={"code": "PROVIDER_ERROR", "message": "Provider could not complete the request."},
|
||||
timestamp=utc_now(),
|
||||
)
|
||||
done = ModelEvent(event=ModelEventType.done, sequence=1, timestamp=utc_now())
|
||||
done = ModelEvent(
|
||||
event=ModelEventType.done, sequence=sequence + 1,
|
||||
data={"status": "failed"}, timestamp=utc_now()
|
||||
)
|
||||
yield as_sse(error.event.value, error.model_dump_json())
|
||||
yield as_sse(done.event.value, done.model_dump_json())
|
||||
|
||||
@@ -436,9 +491,7 @@ async def list_skills() -> SkillListResponse:
|
||||
return SkillListResponse(items=container.skills.list())
|
||||
|
||||
|
||||
@router.get(
|
||||
"/skills/{skill_id}", response_model=Skill, tags=["Skills"]
|
||||
)
|
||||
@router.get("/skills/{skill_id}", response_model=Skill, tags=["Skills"])
|
||||
async def get_skill(skill_id: str) -> Skill:
|
||||
return extension_call(lambda: container.skills.get(skill_id))
|
||||
|
||||
@@ -478,7 +531,120 @@ async def disable_skill(skill_id: str) -> Skill:
|
||||
)
|
||||
async def uninstall_skill(skill_id: str) -> OperationResponse:
|
||||
extension_call(lambda: container.skills.uninstall(skill_id))
|
||||
return OperationResponse(status="completed", resource_id=skill_id, message="uninstalled")
|
||||
return OperationResponse(
|
||||
status="completed", resource_id=skill_id, message="uninstalled"
|
||||
)
|
||||
|
||||
|
||||
# Independent MCP Server Registry
|
||||
@router.get("/mcp/servers", response_model=McpServerListResponse, tags=["MCP Servers"])
|
||||
async def list_mcp_servers() -> McpServerListResponse:
|
||||
return McpServerListResponse(items=await mcp_call_async(container.mcp_servers.list))
|
||||
|
||||
|
||||
@router.post(
|
||||
"/mcp/servers", response_model=McpServer, status_code=201, tags=["MCP Servers"]
|
||||
)
|
||||
async def create_mcp_server(request: McpServerCreateRequest) -> McpServer:
|
||||
return await mcp_call_async(lambda: container.mcp_servers.create(request))
|
||||
|
||||
|
||||
@router.get("/mcp/servers/{server_id}", response_model=McpServer, tags=["MCP Servers"])
|
||||
async def get_mcp_server(server_id: str) -> McpServer:
|
||||
return await mcp_call_async(lambda: container.mcp_servers.get(server_id))
|
||||
|
||||
|
||||
@router.get(
|
||||
"/mcp/servers/{server_id}/tools",
|
||||
response_model=McpToolSummaryListResponse,
|
||||
tags=["MCP Servers"],
|
||||
)
|
||||
async def list_mcp_server_tools(server_id: str) -> McpToolSummaryListResponse:
|
||||
return McpToolSummaryListResponse(
|
||||
items=await mcp_call_async(lambda: container.mcp_servers.list_tools(server_id))
|
||||
)
|
||||
|
||||
|
||||
@router.put("/mcp/servers/{server_id}", response_model=McpServer, tags=["MCP Servers"])
|
||||
async def update_mcp_server(
|
||||
server_id: str, request: McpServerUpdateRequest
|
||||
) -> McpServer:
|
||||
return await mcp_call_async(
|
||||
lambda: container.mcp_servers.update(server_id, request)
|
||||
)
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/mcp/servers/{server_id}", response_model=OperationResponse, tags=["MCP Servers"]
|
||||
)
|
||||
async def delete_mcp_server(server_id: str) -> OperationResponse:
|
||||
await mcp_call_async(lambda: container.mcp_servers.delete(server_id))
|
||||
return OperationResponse(
|
||||
status="completed", resource_id=server_id, message="deleted"
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/mcp/servers/{server_id}/trust", response_model=McpServer, tags=["MCP Servers"]
|
||||
)
|
||||
async def trust_mcp_server(server_id: str, request: McpServerTrustRequest) -> McpServer:
|
||||
return await mcp_call_async(
|
||||
lambda: container.mcp_servers.trust(server_id, request.command_digest)
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/mcp/servers/{server_id}/test", response_model=McpServer, tags=["MCP Servers"]
|
||||
)
|
||||
async def test_mcp_server(server_id: str) -> McpServer:
|
||||
return await mcp_call_async(lambda: container.mcp_servers.test(server_id))
|
||||
|
||||
|
||||
@router.post(
|
||||
"/mcp/servers/{server_id}/enable", response_model=McpServer, tags=["MCP Servers"]
|
||||
)
|
||||
async def enable_mcp_server(server_id: str) -> McpServer:
|
||||
return await mcp_call_async(lambda: container.mcp_servers.enable(server_id))
|
||||
|
||||
|
||||
@router.post(
|
||||
"/mcp/servers/{server_id}/disable", response_model=McpServer, tags=["MCP Servers"]
|
||||
)
|
||||
async def disable_mcp_server(server_id: str) -> McpServer:
|
||||
return await mcp_call_async(lambda: container.mcp_servers.disable(server_id))
|
||||
|
||||
|
||||
@router.put(
|
||||
"/mcp/servers/{server_id}/secrets/{key}",
|
||||
response_model=McpServerSecretStatus,
|
||||
tags=["MCP Servers"],
|
||||
)
|
||||
async def put_mcp_server_secret(
|
||||
server_id: str,
|
||||
key: str,
|
||||
request: McpServerSecretWriteRequest,
|
||||
kind: str = Query(default="environment", pattern="^(environment|header)$"),
|
||||
) -> McpServerSecretStatus:
|
||||
return await mcp_call_async(
|
||||
lambda: container.mcp_servers.put_secret(
|
||||
server_id, key, request.secret.get_secret_value(), kind=kind
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/mcp/servers/{server_id}/secrets/{key}",
|
||||
response_model=McpServerSecretStatus,
|
||||
tags=["MCP Servers"],
|
||||
)
|
||||
async def delete_mcp_server_secret(
|
||||
server_id: str,
|
||||
key: str,
|
||||
kind: str = Query(default="environment", pattern="^(environment|header)$"),
|
||||
) -> McpServerSecretStatus:
|
||||
return await mcp_call_async(
|
||||
lambda: container.mcp_servers.delete_secret(server_id, key, kind=kind)
|
||||
)
|
||||
|
||||
|
||||
# Plugins
|
||||
@@ -570,11 +736,15 @@ async def restart_plugin_host(plugin_id: str) -> OperationResponse:
|
||||
)
|
||||
async def uninstall_plugin(plugin_id: str) -> OperationResponse:
|
||||
plugin = extension_call(lambda: container.plugins.get(plugin_id))
|
||||
dependent_skills = container.skills.depending_on_tools(plugin.manifest.contributes.tools)
|
||||
dependent_skills = container.skills.depending_on_tools(
|
||||
plugin.manifest.contributes.tools
|
||||
)
|
||||
await extension_call_async(
|
||||
lambda: container.plugins.uninstall(plugin_id, dependent_skills)
|
||||
)
|
||||
return OperationResponse(status="completed", resource_id=plugin_id, message="uninstalled")
|
||||
return OperationResponse(
|
||||
status="completed", resource_id=plugin_id, message="uninstalled"
|
||||
)
|
||||
|
||||
|
||||
# Plugin Command / Settings Contributions
|
||||
@@ -649,9 +819,7 @@ async def put_plugin_setting_secret(
|
||||
response_model=PluginSecretStatus,
|
||||
tags=["Plugins"],
|
||||
)
|
||||
async def delete_plugin_setting_secret(
|
||||
plugin_id: str, key: str
|
||||
) -> PluginSecretStatus:
|
||||
async def delete_plugin_setting_secret(plugin_id: str, key: str) -> PluginSecretStatus:
|
||||
return extension_call(
|
||||
lambda: container.plugins.delete_setting_secret(plugin_id, key)
|
||||
)
|
||||
@@ -764,15 +932,17 @@ async def update_provider(
|
||||
) -> ProviderConfig:
|
||||
current = configurable_provider_or_404(provider_id).config
|
||||
if provider_id == "mock":
|
||||
raise ApiError(409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be modified.")
|
||||
raise ApiError(
|
||||
409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be modified."
|
||||
)
|
||||
fields = request.model_fields_set
|
||||
if ("name" in fields and request.name is None) or (
|
||||
if ("provider_type" in fields and request.provider_type is None) or ("name" in fields and request.name is None) or (
|
||||
"enabled" in fields and request.enabled is None
|
||||
):
|
||||
raise ApiError(
|
||||
422,
|
||||
"VALIDATION_ERROR",
|
||||
"name and enabled cannot be null when explicitly provided.",
|
||||
"provider_type, name and enabled cannot be null when explicitly provided.",
|
||||
)
|
||||
updates = {name: getattr(request, name) for name in fields}
|
||||
if "credential_id" in fields:
|
||||
@@ -780,7 +950,11 @@ async def update_provider(
|
||||
config = ProviderConfig.model_validate(
|
||||
{**current.model_dump(mode="python"), **updates}
|
||||
)
|
||||
config.capabilities = container.provider_factory.capabilities(config.provider_type)
|
||||
try:
|
||||
adapter = container.provider_factory.build(config)
|
||||
except UnsupportedProviderError as exc:
|
||||
raise ApiError(422, "PROVIDER_TYPE_UNSUPPORTED", "Provider adapter is not supported.") from exc
|
||||
container.providers.replace(config, adapter)
|
||||
return config
|
||||
|
||||
@@ -793,7 +967,11 @@ async def update_provider(
|
||||
async def delete_provider(provider_id: str) -> OperationResponse:
|
||||
configurable_provider_or_404(provider_id)
|
||||
if provider_id == "mock":
|
||||
raise ApiError(409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be deleted.")
|
||||
raise ApiError(
|
||||
409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be deleted."
|
||||
)
|
||||
if container.model_routing.uses_provider(provider_id):
|
||||
raise ApiError(409, "PROVIDER_IN_USE", "请先在索引与模型中解除该提供商的模型绑定。")
|
||||
container.providers.unregister(provider_id)
|
||||
return OperationResponse(status="completed", resource_id=provider_id)
|
||||
|
||||
@@ -871,7 +1049,9 @@ async def create_task(request: TaskCreateRequest) -> Task:
|
||||
async def get_task(task_id: str) -> Task:
|
||||
task = task_service.get_task(task_id)
|
||||
if task is None:
|
||||
raise ApiError(404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id})
|
||||
raise ApiError(
|
||||
404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id}
|
||||
)
|
||||
return task
|
||||
|
||||
|
||||
@@ -887,11 +1067,35 @@ async def update_task(task_id: str, request: TaskUpdateRequest) -> Task:
|
||||
)
|
||||
async def delete_task(task_id: str) -> OperationResponse:
|
||||
if not task_service.delete_task(task_id):
|
||||
raise ApiError(404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id})
|
||||
raise ApiError(
|
||||
404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id}
|
||||
)
|
||||
return OperationResponse(status="completed", resource_id=task_id, message="deleted")
|
||||
|
||||
|
||||
# Media and index
|
||||
@router.get("/model-routing", response_model=ModelRoutingResponse, tags=["Providers"])
|
||||
async def get_model_routing() -> ModelRoutingResponse:
|
||||
return container.model_routing.describe()
|
||||
|
||||
|
||||
@router.put("/model-routing", response_model=ModelRoutingResponse, tags=["Providers"])
|
||||
async def update_model_routing(request: ModelRoutingConfig) -> ModelRoutingResponse:
|
||||
return container.model_routing.update(request)
|
||||
|
||||
|
||||
@router.post("/models/embeddings", response_model=EmbeddingResult, tags=["Providers"])
|
||||
async def create_embeddings(request: EmbeddingRequest) -> EmbeddingResult:
|
||||
return await container.model_routing.embed(request.texts)
|
||||
|
||||
|
||||
@router.post("/media/speaker-matches", response_model=SpeakerMatchResult, tags=["Media"])
|
||||
async def match_speakers(request: SpeakerMatchRequest) -> SpeakerMatchResult:
|
||||
return await container.model_routing.match_speakers(
|
||||
attachment_path(request.attachment_id), attachment_path(request.reference_attachment_id),
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/media/transcriptions",
|
||||
response_model=TranscriptionJob,
|
||||
@@ -899,8 +1103,8 @@ async def delete_task(task_id: str) -> OperationResponse:
|
||||
tags=["Media"],
|
||||
)
|
||||
async def create_transcription(request: TranscriptionRequest) -> TranscriptionJob:
|
||||
return transcription_service.create_transcription(
|
||||
request.attachment_id, request.language
|
||||
return await transcription_service.create_transcription(
|
||||
request.attachment_id, request.language, diarization=request.diarization
|
||||
)
|
||||
|
||||
|
||||
@@ -937,5 +1141,172 @@ async def rebuild_index(request: IndexRebuildRequest) -> IndexJob:
|
||||
async def get_index_job(job_id: str) -> IndexJob:
|
||||
job = index_service.get_job(job_id)
|
||||
if job is None:
|
||||
raise ApiError(404, "RESOURCE_NOT_FOUND", "index job not found", {"job_id": job_id})
|
||||
raise ApiError(
|
||||
404, "RESOURCE_NOT_FOUND", "index job not found", {"job_id": job_id}
|
||||
)
|
||||
return job
|
||||
|
||||
|
||||
# Benchmark
|
||||
@router.get(
|
||||
"/benchmarks/datasets",
|
||||
response_model=BenchmarkDatasetListResponse,
|
||||
tags=["Benchmark"],
|
||||
)
|
||||
async def list_benchmark_datasets(
|
||||
kind: BenchmarkKind = Query(default=BenchmarkKind.rag),
|
||||
) -> BenchmarkDatasetListResponse:
|
||||
return BenchmarkDatasetListResponse(items=benchmark_datasets.list_datasets(kind))
|
||||
|
||||
|
||||
@router.post(
|
||||
"/benchmarks/rag/runs",
|
||||
response_model=BenchmarkRun,
|
||||
status_code=202,
|
||||
tags=["Benchmark"],
|
||||
)
|
||||
async def create_rag_benchmark(request: RAGRunRequest) -> BenchmarkRun:
|
||||
return await benchmark_service.create_rag_run(request)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/benchmarks/runs",
|
||||
response_model=BenchmarkRunListResponse,
|
||||
tags=["Benchmark"],
|
||||
)
|
||||
async def list_benchmark_runs(
|
||||
kind: BenchmarkKind | None = Query(default=None),
|
||||
status: BenchmarkStatus | None = Query(default=None),
|
||||
limit: int = Query(default=50, ge=1, le=100),
|
||||
offset: int = Query(default=0, ge=0),
|
||||
) -> BenchmarkRunListResponse:
|
||||
items, total = benchmark_service.list_runs(
|
||||
kind=kind, status=status, limit=limit, offset=offset
|
||||
)
|
||||
return BenchmarkRunListResponse(
|
||||
items=items, page=PageMeta(total=total, limit=limit, offset=offset)
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/benchmarks/runs/{run_id}",
|
||||
response_model=BenchmarkRun,
|
||||
tags=["Benchmark"],
|
||||
)
|
||||
async def get_benchmark_run(run_id: str) -> BenchmarkRun:
|
||||
run = benchmark_service.get_run(run_id)
|
||||
if run is None:
|
||||
raise ApiError(
|
||||
404, "BENCHMARK_RUN_NOT_FOUND", "benchmark run not found", {"run_id": run_id}
|
||||
)
|
||||
return run
|
||||
|
||||
|
||||
@router.post(
|
||||
"/benchmarks/runs/{run_id}/cancel",
|
||||
response_model=OperationResponse,
|
||||
tags=["Benchmark"],
|
||||
)
|
||||
async def cancel_benchmark_run(run_id: str) -> OperationResponse:
|
||||
run = benchmark_service.cancel_run(run_id)
|
||||
if run is None:
|
||||
raise ApiError(
|
||||
404, "BENCHMARK_RUN_NOT_FOUND", "benchmark run not found", {"run_id": run_id}
|
||||
)
|
||||
return OperationResponse(
|
||||
status="accepted",
|
||||
resource_id=run_id,
|
||||
message=f"Benchmark run status: {run.status.value}",
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/benchmarks/runs/{run_id}/events",
|
||||
response_class=StreamingResponse,
|
||||
responses={
|
||||
200: {
|
||||
"description": "BenchmarkEvent Server-Sent Events stream",
|
||||
"content": {"text/event-stream": {}},
|
||||
}
|
||||
},
|
||||
tags=["Benchmark"],
|
||||
)
|
||||
async def benchmark_events(
|
||||
run_id: str,
|
||||
after_sequence: int = Query(default=-1, ge=-1),
|
||||
last_event_id: str | None = Header(default=None, alias="Last-Event-ID"),
|
||||
) -> StreamingResponse:
|
||||
if benchmark_service.get_run(run_id) is None:
|
||||
raise ApiError(
|
||||
404, "BENCHMARK_RUN_NOT_FOUND", "benchmark run not found", {"run_id": run_id}
|
||||
)
|
||||
|
||||
# SSE 断线重连:Last-Event-ID 优先于 after_sequence,用于从上次收到的事件继续
|
||||
cursor = after_sequence
|
||||
if last_event_id is not None:
|
||||
try:
|
||||
cursor = int(last_event_id)
|
||||
except ValueError as exc:
|
||||
raise ApiError(
|
||||
400,
|
||||
"BENCHMARK_EVENT_CURSOR_INVALID",
|
||||
"Last-Event-ID must be an integer sequence.",
|
||||
{"last_event_id": last_event_id},
|
||||
) from exc
|
||||
if cursor < -1:
|
||||
raise ApiError(
|
||||
400,
|
||||
"BENCHMARK_EVENT_CURSOR_INVALID",
|
||||
"Last-Event-ID must be greater than or equal to -1.",
|
||||
)
|
||||
|
||||
async def stream() -> AsyncIterator[str]:
|
||||
# 先订阅(保证订阅之后产生的事件也能收到),再回放历史事件,最后实时输出新事件
|
||||
terminal = (
|
||||
BenchmarkEventType.run_completed,
|
||||
BenchmarkEventType.run_failed,
|
||||
BenchmarkEventType.run_cancelled,
|
||||
)
|
||||
queue = benchmark_service.subscribe(run_id)
|
||||
try:
|
||||
last_sequence = cursor
|
||||
# 回放按订阅时刻的快照长度遍历,避免列表在回放期间被追加;终止事件同样要结束流,
|
||||
# 防止回放完成后进入实时队列却因序号去重跳过同一终止事件而永久等待。
|
||||
history = benchmark_service.get_events(run_id)
|
||||
for index in range(len(history)):
|
||||
event = history[index]
|
||||
if event.sequence <= cursor:
|
||||
continue
|
||||
yield as_sse(event.event.value, event.model_dump_json(), event_id=event.sequence)
|
||||
last_sequence = event.sequence
|
||||
if event.event in terminal:
|
||||
return
|
||||
if queue is None:
|
||||
return
|
||||
while True:
|
||||
event = await queue.get()
|
||||
if event.sequence <= last_sequence:
|
||||
continue
|
||||
yield as_sse(event.event.value, event.model_dump_json(), event_id=event.sequence)
|
||||
last_sequence = event.sequence
|
||||
if event.event in terminal:
|
||||
return
|
||||
finally:
|
||||
if queue is not None:
|
||||
benchmark_service.unsubscribe(run_id, queue)
|
||||
|
||||
return StreamingResponse(stream(), media_type="text/event-stream")
|
||||
|
||||
|
||||
@router.get(
|
||||
"/benchmarks/runs/{run_id}/report",
|
||||
response_model=BenchmarkReport,
|
||||
tags=["Benchmark"],
|
||||
)
|
||||
async def get_benchmark_report(run_id: str) -> BenchmarkReport:
|
||||
report = benchmark_service.get_report(run_id)
|
||||
if report is None:
|
||||
raise ApiError(
|
||||
404, "BENCHMARK_RUN_NOT_FOUND", "benchmark report not found", {"run_id": run_id}
|
||||
)
|
||||
return report
|
||||
|
||||
@@ -6,7 +6,6 @@ MVP 阶段重建是同步的(数据量小),完成后直接返回 completed
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import shutil
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from uuid import uuid4
|
||||
@@ -16,8 +15,8 @@ from app.config import get_settings
|
||||
from app.contracts import IndexJob, IndexRebuildRequest, IndexStatus
|
||||
from app.errors import ApiError
|
||||
from app.knowledge.parser import parse_note
|
||||
from app.services.note_service import index_note
|
||||
from app.services import task_service
|
||||
from app.services.note_service import index_note, prepare_note_index
|
||||
from app.database.db import connect, transaction
|
||||
from app.services.coordination import serialized_vault_mutation
|
||||
from app.retrieval.vectorstore import SqliteVecStore
|
||||
|
||||
@@ -74,18 +73,7 @@ async def rebuild(request: IndexRebuildRequest) -> IndexJob:
|
||||
{"scope": request.scope, "note_ids": request.note_ids},
|
||||
)
|
||||
|
||||
# 先扫描到内存(失败不会清旧索引),再快照旧库用于失败回滚
|
||||
docs = _scan_vault()
|
||||
settings = get_settings()
|
||||
database_existed = settings.db_path.exists()
|
||||
task_note_links = task_service.note_links() if database_existed else {}
|
||||
backup_path = (
|
||||
settings.db_path.with_name(f"{settings.db_path.name}.{job_id}.bak")
|
||||
if database_existed
|
||||
else None
|
||||
)
|
||||
if backup_path is not None:
|
||||
shutil.copy2(settings.db_path, backup_path)
|
||||
|
||||
_active_job_id = job_id
|
||||
_last_error = None
|
||||
@@ -94,21 +82,34 @@ async def rebuild(request: IndexRebuildRequest) -> IndexJob:
|
||||
created_at=datetime.now(timezone.utc),
|
||||
))
|
||||
try:
|
||||
repository.clear_all()
|
||||
await vector_store.clear()
|
||||
prepared_notes = []
|
||||
for rel, folder, markdown, created, updated in docs:
|
||||
parsed = parse_note(
|
||||
markdown=markdown, file_path=rel, folder=folder, tags=None,
|
||||
created_at=created, updated_at=updated,
|
||||
)
|
||||
await index_note(parsed)
|
||||
task_service.restore_note_links(task_note_links)
|
||||
prepared_notes.append((parsed, await prepare_note_index(parsed)))
|
||||
# All network/model awaits precede the transaction. The concrete SQLite
|
||||
# methods below complete synchronously despite their async interfaces.
|
||||
conn = connect()
|
||||
try:
|
||||
with transaction(conn):
|
||||
task_note_links = dict(conn.execute(
|
||||
"SELECT task_id, note_id FROM tasks WHERE note_id IS NOT NULL"
|
||||
).fetchall())
|
||||
repository.clear_all(conn=conn)
|
||||
await vector_store.clear(conn=conn)
|
||||
for parsed, prepared in prepared_notes:
|
||||
await index_note(parsed, prepared=prepared, conn=conn)
|
||||
for task_id, note_id in task_note_links.items():
|
||||
conn.execute(
|
||||
"UPDATE tasks SET note_id = ? WHERE task_id = ? "
|
||||
"AND EXISTS (SELECT 1 FROM notes WHERE note_id = ?)",
|
||||
(note_id, task_id, note_id),
|
||||
)
|
||||
finally:
|
||||
conn.close()
|
||||
except BaseException as exc:
|
||||
# 重建失败:恢复旧索引,避免留下半成品;记录 failed 任务后向上抛
|
||||
if backup_path is not None and backup_path.exists():
|
||||
shutil.copy2(backup_path, settings.db_path)
|
||||
elif not database_existed:
|
||||
settings.db_path.unlink(missing_ok=True)
|
||||
_remember_job(IndexJob(
|
||||
job_id=job_id, status="failed", scope=request.scope,
|
||||
created_at=datetime.now(timezone.utc),
|
||||
@@ -117,8 +118,6 @@ async def rebuild(request: IndexRebuildRequest) -> IndexJob:
|
||||
raise
|
||||
finally:
|
||||
_active_job_id = None
|
||||
if backup_path is not None:
|
||||
backup_path.unlink(missing_ok=True)
|
||||
|
||||
job = IndexJob(job_id=job_id, status="completed", scope=request.scope, created_at=datetime.now(timezone.utc))
|
||||
_remember_job(job)
|
||||
|
||||
@@ -6,6 +6,8 @@ Markdown 文件是笔记正文的持久化载体(Vault),SQLite/FTS5/向量
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlite3
|
||||
from contextlib import nullcontext
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from uuid import uuid4
|
||||
@@ -16,6 +18,7 @@ from app.database.db import connect, transaction
|
||||
from app.errors import ApiError
|
||||
from app.knowledge.parser import ParsedNote, parse_note
|
||||
from app.retrieval.embedding import HashEmbeddingProvider
|
||||
from app.retrieval import routed_vectors
|
||||
from app.retrieval.vectorstore import SqliteVecStore, VectorRecord
|
||||
from app.services.coordination import serialized_vault_mutation
|
||||
from app.services.vault_paths import (
|
||||
@@ -71,17 +74,34 @@ def _delete_markdown(rel_path: str) -> None:
|
||||
path.unlink()
|
||||
|
||||
|
||||
async def index_note(parsed: ParsedNote) -> None:
|
||||
PreparedIndex = tuple[list[list[float]], routed_vectors.RemoteEmbeddings | None]
|
||||
|
||||
|
||||
async def prepare_note_index(parsed: ParsedNote) -> PreparedIndex:
|
||||
"""Compute vectors before opening a write transaction (including API I/O)."""
|
||||
texts = [block.content for block in parsed.blocks]
|
||||
vectors = await embedding.embed_documents(texts)
|
||||
remote = await routed_vectors.embed_remote(texts)
|
||||
return vectors, remote
|
||||
|
||||
|
||||
async def index_note(
|
||||
parsed: ParsedNote, *, prepared: PreparedIndex | None = None,
|
||||
conn: sqlite3.Connection | None = None,
|
||||
) -> None:
|
||||
"""把解析结果写入元数据 + FTS5 + 向量(三层可重建索引),单事务保证原子性。
|
||||
|
||||
元数据与向量在同一连接、同一事务内提交,避免「新元数据已提交、向量写入失败」的
|
||||
半提交状态。替换元数据时拿到旧 block_id:清理已删除/内容变化的旧向量,只为新增
|
||||
block 写向量(内容未变的 block 其向量仍有效,无需重复写入)。
|
||||
"""
|
||||
vectors = await embedding.embed_documents([block.content for block in parsed.blocks])
|
||||
conn = connect()
|
||||
if conn is not None and prepared is None:
|
||||
raise ValueError("Prepare embeddings before supplying a write connection")
|
||||
vectors, remote = prepared if prepared is not None else await prepare_note_index(parsed)
|
||||
owns = conn is None
|
||||
conn = conn or connect()
|
||||
try:
|
||||
with transaction(conn):
|
||||
with transaction(conn) if owns else nullcontext():
|
||||
old_block_ids = repository.replace_note_metadata(
|
||||
conn=conn,
|
||||
note_id=parsed.note_id,
|
||||
@@ -105,11 +125,13 @@ async def index_note(parsed: ParsedNote) -> None:
|
||||
if block.block_id in missing_ids
|
||||
]
|
||||
await vector_store.upsert(records, conn=conn)
|
||||
routed_vectors.store_remote(conn, [block.block_id for block in parsed.blocks], remote)
|
||||
repository.set_index_meta(
|
||||
{"embedding_model": embedding.model_id, "embedding_dim": str(embedding.dim)},
|
||||
conn=conn,
|
||||
)
|
||||
finally:
|
||||
if owns:
|
||||
conn.close()
|
||||
|
||||
|
||||
|
||||
@@ -1,37 +1,59 @@
|
||||
"""转写适配层;第一阶段消费文本附件或桌面 Host 预生成的旁路文本。"""
|
||||
"""转写作业:API 优先,本地模型回退;保留已有 Host 文本入口。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import OrderedDict
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from uuid import uuid4
|
||||
|
||||
from app.contracts import TranscriptionJob
|
||||
from app.errors import ApiError
|
||||
from app.services.attachment_service import attachment_path
|
||||
|
||||
_jobs: OrderedDict[str, TranscriptionJob] = OrderedDict()
|
||||
MAX_JOBS = 100
|
||||
|
||||
|
||||
def create_transcription(attachment_id: str, language: str | None = None) -> TranscriptionJob:
|
||||
# TODO(ai-core): 第二阶段接入本地 ASR 队列后,保留相同 Job 契约替换此同步降级实现。
|
||||
del language # 预生成 transcript 暂不需要语言识别。
|
||||
async def create_transcription(attachment_id: str, language: str | None = None, *, diarization: bool = False) -> TranscriptionJob:
|
||||
from app.container import container
|
||||
|
||||
source = attachment_path(attachment_id)
|
||||
transcript = source if source.suffix.lower() in {".txt", ".md"} else Path(f"{source}.txt")
|
||||
job = TranscriptionJob(
|
||||
job_id=f"transcription_{uuid4().hex}",
|
||||
attachment_id=attachment_id,
|
||||
status="completed" if transcript.is_file() else "failed",
|
||||
text=transcript.read_text(encoding="utf-8") if transcript.is_file() else None,
|
||||
error_code=None if transcript.is_file() else "TRANSCRIPTION_BACKEND_UNAVAILABLE",
|
||||
error_message=(
|
||||
None
|
||||
if transcript.is_file()
|
||||
else "No host-generated transcript is available; local speech models are phase two."
|
||||
),
|
||||
status="processing",
|
||||
created_at=datetime.now(timezone.utc),
|
||||
)
|
||||
try:
|
||||
if diarization:
|
||||
# Speaker verification and diarization are different capabilities.
|
||||
raise ApiError(501, "DIARIZATION_NOT_IMPLEMENTED", "说话人分离将在阶段 F 接入,当前不能忽略 diarization 请求。")
|
||||
transcript = source if source.suffix.lower() in {".txt", ".md"} else attachment_path(f"{attachment_id}.txt")
|
||||
# A saved transcript remains an explicit import path, never faked ASR.
|
||||
if transcript.is_file() and (source == transcript or container.model_routing.configuration().transcription is None):
|
||||
with transcript.open("rb") as handle:
|
||||
content = handle.read(1024 * 1024 + 1)
|
||||
if len(content) > 1024 * 1024:
|
||||
raise ApiError(413, "TRANSCRIPT_TOO_LARGE", "Transcript exceeds 1 MiB.")
|
||||
job.text = content.decode("utf-8")
|
||||
if not job.text.strip():
|
||||
raise ApiError(422, "TRANSCRIPT_EMPTY", "Transcript is empty.")
|
||||
job.source = "sidecar"
|
||||
else:
|
||||
result = await container.model_routing.transcribe(source, language)
|
||||
job.text = result.text
|
||||
job.source = result.source
|
||||
job.fallback_reason = result.fallback_reason
|
||||
job.status = "completed"
|
||||
except ApiError as exc:
|
||||
job.status = "failed"
|
||||
job.error_code = exc.code
|
||||
job.error_message = exc.message
|
||||
job.fallback_reason = exc.details.get("fallback_reason")
|
||||
except (OSError, UnicodeError):
|
||||
job.status = "failed"
|
||||
job.error_code = "TRANSCRIPT_UNREADABLE"
|
||||
job.error_message = "Transcript could not be read."
|
||||
_jobs[job.job_id] = job
|
||||
while len(_jobs) > MAX_JOBS:
|
||||
_jobs.popitem(last=False)
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
{
|
||||
"dataset_id": "rag-core-v1",
|
||||
"kind": "rag",
|
||||
"version": "1.0.0",
|
||||
"description": "基础中文笔记检索集(对应 backend/data/vault 内置语料,重建索引后即可复现)",
|
||||
"cases": [
|
||||
{
|
||||
"case_id": "rag-vector-sim",
|
||||
"query": "向量数据库如何进行相似度检索",
|
||||
"expected_note_ids": ["note_c1454740a0e55ef5"],
|
||||
"expected_block_ids": ["blk_07c4c6bce0ec4d12", "blk_605fb3593809f224"],
|
||||
"citation_required": true,
|
||||
"tags": ["向量数据库", "检索"]
|
||||
},
|
||||
{
|
||||
"case_id": "rag-python-func",
|
||||
"query": "Python 如何定义函数",
|
||||
"expected_note_ids": ["note_424c3742c6f0e555"],
|
||||
"expected_block_ids": ["blk_45d48cae2fed40fe", "blk_0768d9c25c2ecf07"],
|
||||
"citation_required": true,
|
||||
"tags": ["python"]
|
||||
},
|
||||
{
|
||||
"case_id": "rag-citation",
|
||||
"query": "搜索结果如何定位到原文位置",
|
||||
"expected_note_ids": ["note_0c619caa30b1614c"],
|
||||
"expected_block_ids": ["blk_3f6fcead71c25fc6", "blk_9af7b12e9ce909fc"],
|
||||
"citation_required": true,
|
||||
"tags": ["RAG"]
|
||||
},
|
||||
{
|
||||
"case_id": "rag-hybrid",
|
||||
"query": "混合检索怎么融合全文和向量",
|
||||
"expected_note_ids": ["note_c1454740a0e55ef5"],
|
||||
"expected_block_ids": ["blk_82b45418dba9f720"],
|
||||
"citation_required": true,
|
||||
"tags": ["检索"]
|
||||
},
|
||||
{
|
||||
"case_id": "rag-tech-stack",
|
||||
"query": "这个项目用什么后端和检索技术",
|
||||
"expected_note_ids": ["note_3327e6cf18f3701f"],
|
||||
"expected_block_ids": ["blk_feb2a9c42e7d31ad"],
|
||||
"citation_required": false,
|
||||
"tags": ["项目"]
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -1,26 +1,12 @@
|
||||
import asyncio
|
||||
import threading
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from app.main import health, service_status
|
||||
from app.routes import (
|
||||
get_index_status,
|
||||
list_notes,
|
||||
list_plugins,
|
||||
list_provider_presets,
|
||||
list_providers,
|
||||
list_skills,
|
||||
)
|
||||
from app.routes import (
|
||||
create_provider,
|
||||
create_task,
|
||||
delete_provider,
|
||||
delete_task,
|
||||
get_provider,
|
||||
get_task,
|
||||
list_tasks,
|
||||
update_provider,
|
||||
update_task,
|
||||
)
|
||||
from app.contracts import (
|
||||
McpServerSecretStatus,
|
||||
McpServerSecretWriteRequest,
|
||||
ProviderCreateRequest,
|
||||
ProviderType,
|
||||
ProviderUpdateRequest,
|
||||
@@ -28,6 +14,61 @@ from app.contracts import (
|
||||
TaskStatus,
|
||||
TaskUpdateRequest,
|
||||
)
|
||||
from app.main import health, service_status
|
||||
from app.routes import (
|
||||
create_provider,
|
||||
create_task,
|
||||
delete_provider,
|
||||
delete_task,
|
||||
get_index_status,
|
||||
get_provider,
|
||||
get_task,
|
||||
list_notes,
|
||||
list_plugins,
|
||||
list_provider_presets,
|
||||
list_providers,
|
||||
list_skills,
|
||||
list_tasks,
|
||||
update_provider,
|
||||
update_task,
|
||||
)
|
||||
|
||||
|
||||
def test_mcp_secret_routes_offload_blocking_lifecycle_work(monkeypatch) -> None:
|
||||
from app import routes
|
||||
|
||||
caller_thread = threading.get_ident()
|
||||
worker_threads: list[int] = []
|
||||
|
||||
class FakeMcpRegistry:
|
||||
def put_secret(self, server_id, key, secret, *, kind):
|
||||
worker_threads.append(threading.get_ident())
|
||||
return McpServerSecretStatus(key=key, configured=True)
|
||||
|
||||
def delete_secret(self, server_id, key, *, kind):
|
||||
worker_threads.append(threading.get_ident())
|
||||
return McpServerSecretStatus(key=key, configured=False)
|
||||
|
||||
monkeypatch.setattr(
|
||||
routes,
|
||||
"container",
|
||||
SimpleNamespace(mcp_servers=FakeMcpRegistry()),
|
||||
)
|
||||
written = asyncio.run(
|
||||
routes.put_mcp_server_secret(
|
||||
"server-1",
|
||||
"TOKEN",
|
||||
McpServerSecretWriteRequest(secret="hidden"),
|
||||
kind="environment",
|
||||
)
|
||||
)
|
||||
deleted = asyncio.run(
|
||||
routes.delete_mcp_server_secret("server-1", "TOKEN", kind="environment")
|
||||
)
|
||||
|
||||
assert written.configured is True
|
||||
assert deleted.configured is False
|
||||
assert worker_threads and all(item != caller_thread for item in worker_threads)
|
||||
|
||||
|
||||
def test_health() -> None:
|
||||
@@ -36,6 +77,172 @@ def test_health() -> None:
|
||||
assert response.model_dump() == {"status": "ok"}
|
||||
|
||||
|
||||
def test_mcp_create_and_trust_are_not_executed_on_event_loop(monkeypatch) -> None:
|
||||
from app import routes
|
||||
from app.contracts import McpServerCreateRequest, McpServerTrustRequest
|
||||
|
||||
caller = threading.get_ident()
|
||||
workers = []
|
||||
|
||||
class Registry:
|
||||
def create(self, request):
|
||||
workers.append(threading.get_ident())
|
||||
return "created"
|
||||
|
||||
def trust(self, server_id, digest):
|
||||
workers.append(threading.get_ident())
|
||||
return "trusted"
|
||||
|
||||
monkeypatch.setattr(routes, "container", SimpleNamespace(mcp_servers=Registry()))
|
||||
assert (
|
||||
asyncio.run(
|
||||
routes.create_mcp_server(McpServerCreateRequest(name="test", command="uvx"))
|
||||
)
|
||||
== "created"
|
||||
)
|
||||
assert (
|
||||
asyncio.run(
|
||||
routes.trust_mcp_server(
|
||||
"test", McpServerTrustRequest(command_digest="a" * 64)
|
||||
)
|
||||
)
|
||||
== "trusted"
|
||||
)
|
||||
assert len(workers) == 2
|
||||
assert all(worker != caller for worker in workers)
|
||||
|
||||
|
||||
def test_mcp_split_config_and_secret_requests_persist_without_plaintext(
|
||||
monkeypatch,
|
||||
) -> None:
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from app import routes
|
||||
from app.agent.tools import ToolRegistry
|
||||
from app.config import get_settings
|
||||
from app.extensions.mcp_registry import McpServerRegistry
|
||||
from app.main import app
|
||||
from app.providers.credentials import EncryptedCredentialStore
|
||||
|
||||
service = McpServerRegistry(
|
||||
ToolRegistry(),
|
||||
EncryptedCredentialStore(),
|
||||
get_settings().data_dir,
|
||||
allow_process_launch=True,
|
||||
)
|
||||
monkeypatch.setattr(routes, "container", SimpleNamespace(mcp_servers=service))
|
||||
client = TestClient(app)
|
||||
config = {
|
||||
"name": "MiniMax configuration test",
|
||||
"command": "uvx",
|
||||
"environment": {"MINIMAX_API_HOST": "https://api.minimaxi.com"},
|
||||
"secret_environment_keys": ["MINIMAX_API_KEY"],
|
||||
"startup_timeout_seconds": 120,
|
||||
"tool_timeout_seconds": 300,
|
||||
}
|
||||
# Reproduce the old frontend payload. The backend still enforces separation.
|
||||
invalid = client.post(
|
||||
"/api/mcp/servers",
|
||||
json={
|
||||
**config,
|
||||
"environment": {
|
||||
**config["environment"],
|
||||
"MINIMAX_API_KEY": "synthetic-only",
|
||||
},
|
||||
},
|
||||
)
|
||||
assert invalid.status_code == 422
|
||||
assert invalid.json()["error"]["code"] == "MCP_ENVIRONMENT_INVALID"
|
||||
created = client.post("/api/mcp/servers", json=config)
|
||||
assert created.status_code == 201
|
||||
server_id = created.json()["server_id"]
|
||||
saved = client.put(
|
||||
f"/api/mcp/servers/{server_id}/secrets/MINIMAX_API_KEY",
|
||||
json={"secret": "synthetic-only"},
|
||||
)
|
||||
assert saved.status_code == 200
|
||||
current = client.get(f"/api/mcp/servers/{server_id}")
|
||||
assert current.json()["secret_environment"] == {"MINIMAX_API_KEY": True}
|
||||
assert "synthetic-only" not in current.text
|
||||
assert "synthetic-only" not in service._path.read_text(encoding="utf-8")
|
||||
_, credentials_path = service.credentials._paths()
|
||||
assert "synthetic-only" not in credentials_path.read_text(encoding="utf-8")
|
||||
assert not current.json()["enabled"] # Saving never starts a third-party process.
|
||||
client.close()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("operation", ["create", "trust"])
|
||||
def test_mcp_lifecycle_lock_contention_keeps_event_loop_responsive(
|
||||
monkeypatch,
|
||||
operation,
|
||||
) -> None:
|
||||
from app import routes
|
||||
from app.agent.tools import ToolRegistry
|
||||
from app.config import get_settings
|
||||
from app.contracts import McpServerCreateRequest, McpServerTrustRequest
|
||||
from app.extensions.mcp_registry import McpServerRegistry
|
||||
from app.providers.credentials import EncryptedCredentialStore
|
||||
|
||||
service = McpServerRegistry(
|
||||
ToolRegistry(),
|
||||
EncryptedCredentialStore(),
|
||||
get_settings().data_dir,
|
||||
allow_process_launch=True,
|
||||
)
|
||||
request = McpServerCreateRequest(
|
||||
name="Lock contention fixture", command="not-executed"
|
||||
)
|
||||
server = service.create(request)
|
||||
monkeypatch.setattr(routes, "container", SimpleNamespace(mcp_servers=service))
|
||||
entered = threading.Event()
|
||||
locked = threading.Event()
|
||||
release = threading.Event()
|
||||
original = getattr(service, operation)
|
||||
|
||||
def observed(*args):
|
||||
entered.set()
|
||||
return original(*args)
|
||||
|
||||
def hold_lifecycle_lock():
|
||||
with service._lifecycle_lock:
|
||||
locked.set()
|
||||
release.wait(timeout=5)
|
||||
|
||||
monkeypatch.setattr(service, operation, observed)
|
||||
holder = threading.Thread(target=hold_lifecycle_lock, daemon=True)
|
||||
holder.start()
|
||||
# An independent watchdog lets the test fail rather than hang if a regression
|
||||
# blocks the event loop itself (an asyncio timeout alone cannot catch that).
|
||||
watchdog = threading.Timer(5, release.set)
|
||||
watchdog.start()
|
||||
|
||||
async def exercise():
|
||||
pending = asyncio.create_task(
|
||||
routes.create_mcp_server(request)
|
||||
if operation == "create"
|
||||
else routes.trust_mcp_server(
|
||||
server.server_id,
|
||||
McpServerTrustRequest(command_digest=server.command_digest),
|
||||
)
|
||||
)
|
||||
try:
|
||||
assert await asyncio.to_thread(entered.wait, 2)
|
||||
assert not pending.done()
|
||||
assert not release.is_set()
|
||||
assert (await health()).status == "ok"
|
||||
finally:
|
||||
release.set()
|
||||
await pending
|
||||
|
||||
try:
|
||||
assert locked.wait(timeout=2)
|
||||
asyncio.run(exercise())
|
||||
finally:
|
||||
release.set()
|
||||
watchdog.cancel()
|
||||
holder.join(timeout=2)
|
||||
|
||||
|
||||
def test_service_status() -> None:
|
||||
response = asyncio.run(service_status())
|
||||
|
||||
@@ -52,7 +259,9 @@ def test_core_collections_are_typed() -> None:
|
||||
|
||||
assert notes.items == []
|
||||
assert notes.page.limit == 20
|
||||
assert [skill.manifest.skill_id for skill in skills.items] == ["knowledge-assistant"]
|
||||
assert [skill.manifest.skill_id for skill in skills.items] == [
|
||||
"knowledge-assistant"
|
||||
]
|
||||
assert skills.items[0].status == "ready"
|
||||
assert [plugin.manifest.plugin_id for plugin in plugins.items] == ["text-tools"]
|
||||
assert plugins.items[0].status == "ready"
|
||||
@@ -73,9 +282,15 @@ def test_provider_presets_include_openai_and_deepseek() -> None:
|
||||
def test_provider_presets_static_route_precedes_provider_id_route() -> None:
|
||||
from app.routes import router
|
||||
|
||||
get_paths = [route.path for route in router.routes if "GET" in getattr(route, "methods", set())]
|
||||
get_paths = [
|
||||
route.path
|
||||
for route in router.routes
|
||||
if "GET" in getattr(route, "methods", set())
|
||||
]
|
||||
|
||||
assert get_paths.index("/api/providers/presets") < get_paths.index("/api/providers/{provider_id}")
|
||||
assert get_paths.index("/api/providers/presets") < get_paths.index(
|
||||
"/api/providers/{provider_id}"
|
||||
)
|
||||
|
||||
|
||||
def test_openapi_contains_documented_frontend_interfaces() -> None:
|
||||
@@ -101,6 +316,14 @@ def test_openapi_contains_documented_frontend_interfaces() -> None:
|
||||
"/api/plugins/{plugin_id}/settings/{key}/secret",
|
||||
"/api/plugins/{plugin_id}/enable",
|
||||
"/api/plugins/{plugin_id}/disable",
|
||||
"/api/mcp/servers",
|
||||
"/api/mcp/servers/{server_id}",
|
||||
"/api/mcp/servers/{server_id}/tools",
|
||||
"/api/mcp/servers/{server_id}/trust",
|
||||
"/api/mcp/servers/{server_id}/test",
|
||||
"/api/mcp/servers/{server_id}/enable",
|
||||
"/api/mcp/servers/{server_id}/disable",
|
||||
"/api/mcp/servers/{server_id}/secrets/{key}",
|
||||
"/api/providers/test",
|
||||
"/api/providers/presets",
|
||||
"/api/credentials/{credential_id}",
|
||||
|
||||
@@ -0,0 +1,567 @@
|
||||
"""Benchmark 服务的单元与端到端测试。
|
||||
|
||||
沿用 conftest 的隔离机制:APP_DATA_DIR / DB / Vault 都指向临时目录,benchmark
|
||||
数据集也落在临时目录(settings.benchmark_datasets_path),不读写真实数据。
|
||||
|
||||
运行采用「创建即 queued + 后台 Task 执行」的异步模型,测试通过 _run 在同一事件循环内
|
||||
创建并等待后台任务结束,得到终态 BenchmarkRun 后再断言。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from app.benchmarks import datasets, metrics as m, service
|
||||
from app.config import get_settings
|
||||
from app.contracts import (
|
||||
BenchmarkKind,
|
||||
BenchmarkRun,
|
||||
BenchmarkStatus,
|
||||
RAGRunRequest,
|
||||
SearchMode,
|
||||
)
|
||||
from app.errors import ApiError
|
||||
|
||||
|
||||
def _write_dataset(dataset_id: str, cases: list[dict], *, kind: str = "rag") -> None:
|
||||
directory = get_settings().benchmark_datasets_path
|
||||
directory.mkdir(parents=True, exist_ok=True)
|
||||
payload = {
|
||||
"dataset_id": dataset_id,
|
||||
"kind": kind,
|
||||
"version": "1.0.0",
|
||||
"description": "test dataset",
|
||||
"cases": cases,
|
||||
}
|
||||
(directory / f"{dataset_id}.json").write_text(
|
||||
json.dumps(payload, ensure_ascii=False), encoding="utf-8"
|
||||
)
|
||||
|
||||
|
||||
def _write_raw(dataset_id: str, raw: dict) -> None:
|
||||
directory = get_settings().benchmark_datasets_path
|
||||
directory.mkdir(parents=True, exist_ok=True)
|
||||
(directory / f"{dataset_id}.json").write_text(
|
||||
json.dumps(raw, ensure_ascii=False), encoding="utf-8"
|
||||
)
|
||||
|
||||
|
||||
def _run(request: RAGRunRequest):
|
||||
"""创建运行并在同一事件循环内等待后台任务结束,返回终态 BenchmarkRun。"""
|
||||
from app.contracts import BenchmarkRun
|
||||
|
||||
async def _execute() -> BenchmarkRun:
|
||||
run = await service.create_rag_run(request)
|
||||
return await service.wait_for_run(run.run_id)
|
||||
|
||||
return asyncio.run(_execute())
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 指标纯函数
|
||||
# --------------------------------------------------------------------------- #
|
||||
def test_hit_at_k_and_recall() -> None:
|
||||
retrieved = ["a", "b", "c"]
|
||||
expected = {"b", "z"}
|
||||
|
||||
assert m.hit_at_k(retrieved, expected, 1) is False
|
||||
assert m.hit_at_k(retrieved, expected, 2) is True
|
||||
assert m.recall_at_k(retrieved, expected, 5) == 0.5 # 只召回 b
|
||||
|
||||
|
||||
def test_recall_at_k_dedups_duplicate_notes() -> None:
|
||||
# 同一 Note 经多个 Block 重复出现,去重后 Recall 不应超过 1
|
||||
assert m.recall_at_k(["note-a", "note-a"], {"note-a"}, 2) == 1.0
|
||||
assert m.recall_at_k(["note-a", "note-a", "note-b"], {"note-a"}, 3) == 1.0
|
||||
|
||||
|
||||
def test_reciprocal_rank_and_citation_hit() -> None:
|
||||
assert m.reciprocal_rank(["x", "a", "b"], {"b"}) == 1 / 3
|
||||
assert m.reciprocal_rank(["x"], {"b"}) == 0.0
|
||||
assert m.citation_hit(["blk_1"], {"blk_1"}) is True
|
||||
assert m.citation_hit(["blk_2"], {"blk_1"}) is False
|
||||
assert m.citation_hit([], {"blk_1"}) is False
|
||||
|
||||
|
||||
def test_percentile() -> None:
|
||||
assert m.percentile([1.0, 2.0, 3.0, 4.0], 50.0) == 2.5
|
||||
assert m.percentile([], 50.0) == 0.0
|
||||
assert m.percentile([7.0], 95.0) == 7.0
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Dataset 注册与校验
|
||||
# --------------------------------------------------------------------------- #
|
||||
def test_list_datasets_empty_by_default() -> None:
|
||||
assert datasets.list_datasets(BenchmarkKind.rag) == []
|
||||
|
||||
|
||||
def test_load_missing_dataset_raises() -> None:
|
||||
with pytest.raises(ApiError) as exc:
|
||||
datasets.load_dataset("does-not-exist", BenchmarkKind.rag)
|
||||
assert exc.value.status_code == 404
|
||||
assert exc.value.code == "BENCHMARK_DATASET_NOT_FOUND"
|
||||
|
||||
|
||||
def test_dataset_without_expected_ids_is_invalid() -> None:
|
||||
_write_dataset("bad-v1", [{"case_id": "x", "query": "q", "citation_required": False}])
|
||||
with pytest.raises(ApiError) as exc:
|
||||
datasets.load_dataset("bad-v1", BenchmarkKind.rag)
|
||||
assert exc.value.code == "BENCHMARK_DATASET_INVALID"
|
||||
|
||||
|
||||
def test_dataset_kind_mismatch_is_invalid() -> None:
|
||||
_write_dataset("agent-v1", [{"case_id": "x", "query": "q", "expected_note_ids": ["n"]}], kind="agent")
|
||||
with pytest.raises(ApiError) as exc:
|
||||
datasets.load_dataset("agent-v1", BenchmarkKind.rag)
|
||||
assert exc.value.code == "BENCHMARK_DATASET_INVALID"
|
||||
|
||||
|
||||
def test_citation_required_requires_expected_block_ids() -> None:
|
||||
# citation_required=true 却没有 expected_block_ids,无法计算 Citation Hit Rate,应拒绝
|
||||
_write_dataset(
|
||||
"cit-req-v1",
|
||||
[{"case_id": "x", "query": "q", "expected_note_ids": ["n"], "citation_required": True}],
|
||||
)
|
||||
with pytest.raises(ApiError) as exc:
|
||||
datasets.load_dataset("cit-req-v1", BenchmarkKind.rag)
|
||||
assert exc.value.code == "BENCHMARK_DATASET_INVALID"
|
||||
|
||||
|
||||
def test_list_datasets_skips_corrupted_structure() -> None:
|
||||
# 合法 JSON 但字段结构错误(cases: 42),列表接口应隔离该文件而非整体 500
|
||||
_write_raw("bad-structure", {"dataset_id": "bad-structure", "kind": "rag", "cases": 42})
|
||||
_write_dataset("good-v1", [{"case_id": "x", "query": "q", "expected_note_ids": ["n"]}])
|
||||
|
||||
infos = datasets.list_datasets(BenchmarkKind.rag)
|
||||
ids = {info.dataset_id for info in infos}
|
||||
assert "good-v1" in ids
|
||||
assert "bad-structure" not in ids
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 请求校验(空 / 重复 modes)
|
||||
# --------------------------------------------------------------------------- #
|
||||
def test_empty_modes_rejected() -> None:
|
||||
with pytest.raises(ValidationError):
|
||||
RAGRunRequest(dataset_id="x", modes=[])
|
||||
|
||||
|
||||
def test_duplicate_modes_rejected() -> None:
|
||||
with pytest.raises(ValidationError):
|
||||
RAGRunRequest(dataset_id="x", modes=[SearchMode.fts, SearchMode.fts])
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# RAG Benchmark 端到端
|
||||
# --------------------------------------------------------------------------- #
|
||||
def _single_note_case() -> tuple[str, str, dict]:
|
||||
from app.services import note_service
|
||||
|
||||
note = asyncio.run(
|
||||
note_service.create_note(
|
||||
title="向量库",
|
||||
markdown="向量数据库用于存储高维向量并支持近似最近邻检索。",
|
||||
folder="",
|
||||
tags=["向量"],
|
||||
)
|
||||
)
|
||||
case = {
|
||||
"case_id": "c1",
|
||||
"query": "向量数据库相似度检索",
|
||||
"expected_note_ids": [note.note_id],
|
||||
"expected_block_ids": [note.blocks[0].block_id],
|
||||
"citation_required": True,
|
||||
"tags": ["向量"],
|
||||
}
|
||||
return note.note_id, note.blocks[0].block_id, case
|
||||
|
||||
|
||||
def test_rag_benchmark_end_to_end() -> None:
|
||||
_, _, case = _single_note_case()
|
||||
_write_dataset("e2e-v1", [case])
|
||||
|
||||
run = _run(RAGRunRequest(dataset_id="e2e-v1", modes=[SearchMode.fts]))
|
||||
|
||||
assert run.status.value == "completed"
|
||||
assert run.dataset_hash.startswith("sha256:")
|
||||
assert run.metrics is not None
|
||||
|
||||
fts = run.metrics["fts"]
|
||||
assert fts["hit_at_1"] == 1.0
|
||||
assert fts["recall_at_k"] == 1.0
|
||||
assert fts["mrr"] == 1.0
|
||||
assert fts["citation_hit_rate"] == 1.0
|
||||
assert fts["p50_latency_ms"] >= 0.0
|
||||
assert fts["p95_latency_ms"] >= fts["p50_latency_ms"]
|
||||
|
||||
|
||||
def test_rag_benchmark_all_modes_produce_metrics() -> None:
|
||||
_, _, case = _single_note_case()
|
||||
_write_dataset("e2e-modes-v1", [case])
|
||||
|
||||
run = _run(RAGRunRequest(dataset_id="e2e-modes-v1"))
|
||||
assert run.status.value == "completed"
|
||||
|
||||
for mode in ("fts", "vector", "hybrid"):
|
||||
assert mode in run.metrics
|
||||
for key in ("hit_at_1", "hit_at_5", "recall_at_k", "mrr", "citation_hit_rate"):
|
||||
assert 0.0 <= run.metrics[mode][key] <= 1.0
|
||||
|
||||
|
||||
def test_config_snapshot_records_index_and_models() -> None:
|
||||
_, _, case = _single_note_case()
|
||||
_write_dataset("snapshot-v1", [case])
|
||||
|
||||
run = _run(RAGRunRequest(dataset_id="snapshot-v1", modes=[SearchMode.fts]))
|
||||
|
||||
snapshot = run.config_snapshot
|
||||
assert snapshot["index_meta"] is not None
|
||||
assert snapshot["embedding"]["policy"] == "per_case"
|
||||
assert snapshot["local_embedding"]["version"]
|
||||
assert snapshot["local_embedding"]["dim"]
|
||||
assert snapshot["reranker"]["version"]
|
||||
assert snapshot["retrieval"]["rrf_k"] == 60
|
||||
|
||||
|
||||
def test_benchmark_report_and_events() -> None:
|
||||
_, _, case = _single_note_case()
|
||||
_write_dataset("report-v1", [case])
|
||||
|
||||
run = _run(RAGRunRequest(dataset_id="report-v1", modes=[SearchMode.fts]))
|
||||
report = service.get_report(run.run_id)
|
||||
events = service.get_events(run.run_id)
|
||||
|
||||
assert report is not None
|
||||
assert report.run_id == run.run_id
|
||||
assert len(report.cases) == 1
|
||||
assert report.cases[0].case_id == "c1"
|
||||
assert report.cases[0].hit_at_1 is True
|
||||
|
||||
assert events, "运行应产生事件"
|
||||
assert events[0].event.value == "RunStarted"
|
||||
assert events[-1].event.value == "RunCompleted"
|
||||
|
||||
|
||||
def test_cancel_completed_run_keeps_status() -> None:
|
||||
_, _, case = _single_note_case()
|
||||
_write_dataset("cancel-v1", [case])
|
||||
|
||||
run = _run(RAGRunRequest(dataset_id="cancel-v1", modes=[SearchMode.fts]))
|
||||
assert run.status.value == "completed"
|
||||
|
||||
cancelled = service.cancel_run(run.run_id)
|
||||
assert cancelled.status.value == "completed" # 已结束,不再变 cancelled
|
||||
|
||||
|
||||
def test_cancel_queued_run_marks_cancelled() -> None:
|
||||
_, _, case = _single_note_case()
|
||||
_write_dataset("cancel-queued-v1", [case])
|
||||
|
||||
async def _scenario():
|
||||
run = await service.create_rag_run(
|
||||
RAGRunRequest(dataset_id="cancel-queued-v1", modes=[SearchMode.fts])
|
||||
)
|
||||
service.cancel_run(run.run_id)
|
||||
return await service.wait_for_run(run.run_id)
|
||||
|
||||
run = asyncio.run(_scenario())
|
||||
assert run.status.value == "cancelled"
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 指标聚合:Citation Hit Rate 只统计 citation_required 样本
|
||||
# --------------------------------------------------------------------------- #
|
||||
def test_citation_hit_rate_only_counts_citation_required() -> None:
|
||||
from app.benchmarks import rag as rag_module
|
||||
from app.contracts import RAGCaseResult
|
||||
|
||||
cases = [
|
||||
RAGCaseResult(
|
||||
case_id="a", mode=SearchMode.fts, repeat=0, latency_ms=1.0,
|
||||
citation_hit=True, citation_applicable=True,
|
||||
),
|
||||
RAGCaseResult(
|
||||
case_id="b", mode=SearchMode.fts, repeat=0, latency_ms=1.0,
|
||||
citation_hit=False, citation_applicable=False,
|
||||
),
|
||||
]
|
||||
metrics = rag_module._aggregate(cases, SearchMode.fts)
|
||||
# 只有 citation_applicable(citation_required=true)的样本计入分母
|
||||
assert metrics.citation_hit_rate == 1.0
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 路由接入
|
||||
# --------------------------------------------------------------------------- #
|
||||
def test_benchmark_routes_wired() -> None:
|
||||
from app import routes
|
||||
|
||||
_, _, case = _single_note_case()
|
||||
_write_dataset("route-v1", [case])
|
||||
|
||||
async def _scenario():
|
||||
listed = await routes.list_benchmark_datasets(BenchmarkKind.rag)
|
||||
assert any(item.dataset_id == "route-v1" for item in listed.items)
|
||||
|
||||
run = await routes.create_rag_benchmark(
|
||||
RAGRunRequest(dataset_id="route-v1", modes=[SearchMode.fts])
|
||||
)
|
||||
assert run.status.value == "queued"
|
||||
return await service.wait_for_run(run.run_id)
|
||||
|
||||
run = asyncio.run(_scenario())
|
||||
assert run.status.value == "completed"
|
||||
|
||||
got = asyncio.run(routes.get_benchmark_run(run.run_id))
|
||||
assert got.run_id == run.run_id
|
||||
|
||||
report = asyncio.run(routes.get_benchmark_report(run.run_id))
|
||||
assert report.cases[0].case_id == "c1"
|
||||
|
||||
|
||||
def test_benchmark_run_not_found_raises() -> None:
|
||||
from app import routes
|
||||
|
||||
with pytest.raises(ApiError) as exc:
|
||||
asyncio.run(routes.get_benchmark_run("benchmark_missing"))
|
||||
assert exc.value.code == "BENCHMARK_RUN_NOT_FOUND"
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 审阅回归:索引兼容 / 容量 / 失败样本 / 取消事件 / 数据集隔离
|
||||
# --------------------------------------------------------------------------- #
|
||||
def test_create_rag_run_requires_built_index() -> None:
|
||||
# 空索引(无已索引 block)会让所有模式得到全 0 指标,应在创建时拒绝而非跑出误导结果
|
||||
_write_dataset("empty-index-v1", [{"case_id": "x", "query": "q", "expected_note_ids": ["n"]}])
|
||||
with pytest.raises(ApiError) as exc:
|
||||
asyncio.run(
|
||||
service.create_rag_run(
|
||||
RAGRunRequest(dataset_id="empty-index-v1", modes=[SearchMode.fts])
|
||||
)
|
||||
)
|
||||
assert exc.value.status_code == 409
|
||||
assert exc.value.code == "BENCHMARK_INDEX_INCOMPATIBLE"
|
||||
|
||||
|
||||
def test_capacity_exceeded_when_all_runs_active(monkeypatch) -> None:
|
||||
# 满容量且全为活动(非终态)run 时,无法淘汰,应拒绝创建而非删掉正在运行的 run
|
||||
_, _, case = _single_note_case()
|
||||
_write_dataset("capacity-v1", [case])
|
||||
|
||||
monkeypatch.setattr(service, "MAX_RUNS", 1)
|
||||
fake_id = "benchmark_fake_active"
|
||||
service._runs[fake_id] = BenchmarkRun(
|
||||
run_id=fake_id,
|
||||
kind=BenchmarkKind.rag,
|
||||
dataset_id="capacity-v1",
|
||||
dataset_hash="sha256:fake",
|
||||
status=BenchmarkStatus.queued,
|
||||
created_at=service._now(),
|
||||
)
|
||||
try:
|
||||
with pytest.raises(ApiError) as exc:
|
||||
asyncio.run(
|
||||
service.create_rag_run(
|
||||
RAGRunRequest(dataset_id="capacity-v1", modes=[SearchMode.fts])
|
||||
)
|
||||
)
|
||||
assert exc.value.status_code == 429
|
||||
assert exc.value.code == "BENCHMARK_CAPACITY_EXCEEDED"
|
||||
finally:
|
||||
service._runs.pop(fake_id, None)
|
||||
|
||||
|
||||
def test_failed_samples_counted_as_zero_in_aggregate() -> None:
|
||||
from app.benchmarks import rag as rag_module
|
||||
from app.contracts import RAGCaseResult
|
||||
|
||||
cases = [
|
||||
RAGCaseResult(
|
||||
case_id="ok", mode=SearchMode.fts, repeat=0, latency_ms=10.0,
|
||||
hit_at_1=True, recall=1.0, reciprocal_rank=1.0,
|
||||
citation_hit=True, citation_applicable=True,
|
||||
),
|
||||
RAGCaseResult(
|
||||
case_id="boom", mode=SearchMode.fts, repeat=0, latency_ms=0.0,
|
||||
error="RAG case evaluation failed.",
|
||||
error_code="BENCHMARK_CASE_EVALUATION_FAILED",
|
||||
),
|
||||
]
|
||||
metrics = rag_module._aggregate(cases, SearchMode.fts)
|
||||
|
||||
assert metrics.total_cases == 2
|
||||
assert metrics.successful_cases == 1
|
||||
assert metrics.failed_cases == 1
|
||||
assert metrics.failure_rate == 0.5
|
||||
# 失败样本按零分计入质量指标分母,汇总不虚高
|
||||
assert metrics.hit_at_1 == 0.5
|
||||
assert metrics.recall_at_k == 0.5
|
||||
# 延迟只统计成功样本
|
||||
assert metrics.p50_latency_ms == 10.0
|
||||
|
||||
|
||||
def test_cancel_emits_run_cancelled_event() -> None:
|
||||
_, _, case = _single_note_case()
|
||||
_write_dataset("cancel-event-v1", [case])
|
||||
|
||||
async def _scenario():
|
||||
run = await service.create_rag_run(
|
||||
RAGRunRequest(dataset_id="cancel-event-v1", modes=[SearchMode.fts])
|
||||
)
|
||||
service.cancel_run(run.run_id)
|
||||
return await service.wait_for_run(run.run_id)
|
||||
|
||||
run = asyncio.run(_scenario())
|
||||
assert run.status.value == "cancelled"
|
||||
events = service.get_events(run.run_id)
|
||||
assert events[-1].event.value == "RunCancelled"
|
||||
|
||||
|
||||
def test_load_dataset_ignores_corrupted_unrelated_files() -> None:
|
||||
# 无关文件损坏(非法 JSON / 顶层非对象)不应阻断目标数据集加载
|
||||
directory = get_settings().benchmark_datasets_path
|
||||
directory.mkdir(parents=True, exist_ok=True)
|
||||
(directory / "broken.json").write_text("{ not valid json", encoding="utf-8")
|
||||
(directory / "array.json").write_text('["a", "b"]', encoding="utf-8")
|
||||
_write_dataset("ok-v1", [{"case_id": "x", "query": "q", "expected_note_ids": ["n"]}])
|
||||
|
||||
dataset = datasets.load_dataset("ok-v1", BenchmarkKind.rag)
|
||||
assert dataset.dataset_id == "ok-v1"
|
||||
assert len(dataset.cases) == 1
|
||||
|
||||
|
||||
def test_load_dataset_top_level_must_be_object() -> None:
|
||||
_write_raw("array-top", ["a", "b"])
|
||||
with pytest.raises(ApiError) as exc:
|
||||
datasets.load_dataset("array-top", BenchmarkKind.rag)
|
||||
assert exc.value.code == "BENCHMARK_DATASET_INVALID"
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 审阅回归:运行中取消 / 仅块标注 / SSE 终止事件
|
||||
# --------------------------------------------------------------------------- #
|
||||
def test_cancel_running_benchmark_stops_early() -> None:
|
||||
"""运行中取消应在样本边界及时生效,而非跑完全部样本(审阅 P1)。"""
|
||||
from app.benchmarks import service
|
||||
from app.services import note_service
|
||||
|
||||
note = asyncio.run(
|
||||
note_service.create_note(
|
||||
title="取消回归", markdown="向量数据库用于存储高维向量。", folder="", tags=["向量"]
|
||||
)
|
||||
)
|
||||
cases = [
|
||||
{
|
||||
"case_id": f"c{i}",
|
||||
"query": "向量数据库",
|
||||
"expected_note_ids": [note.note_id],
|
||||
"expected_block_ids": [note.blocks[0].block_id],
|
||||
"citation_required": True,
|
||||
}
|
||||
for i in range(50)
|
||||
]
|
||||
_write_dataset("cancel-running-v1", cases)
|
||||
|
||||
async def _scenario():
|
||||
run = await service.create_rag_run(
|
||||
RAGRunRequest(dataset_id="cancel-running-v1", modes=[SearchMode.fts])
|
||||
)
|
||||
|
||||
async def _cancel_after_start():
|
||||
# 取消通过事件循环调度(独立 Task),而非同步直调,才能复现事件循环饥饿
|
||||
while service.get_run(run.run_id).status == BenchmarkStatus.queued:
|
||||
await asyncio.sleep(0)
|
||||
service.cancel_run(run.run_id)
|
||||
|
||||
cancel_task = asyncio.create_task(_cancel_after_start())
|
||||
finished = await service.wait_for_run(run.run_id)
|
||||
await cancel_task
|
||||
return finished
|
||||
|
||||
run = asyncio.run(_scenario())
|
||||
assert run.status.value == "cancelled"
|
||||
completed = sum(
|
||||
1 for e in service.get_events(run.run_id) if e.event.value == "CaseCompleted"
|
||||
)
|
||||
assert completed < 50 # 未跑完全部样本,证明取消在样本边界生效
|
||||
|
||||
|
||||
def test_block_only_annotation_resolves_note_and_scores() -> None:
|
||||
"""仅标注 expected_block_ids 的样本应按块反查笔记评分,而非零分(审阅 P2)。"""
|
||||
from app.services import note_service
|
||||
|
||||
note = asyncio.run(
|
||||
note_service.create_note(
|
||||
title="仅块标注", markdown="向量数据库存储高维向量。", folder="", tags=["向量"]
|
||||
)
|
||||
)
|
||||
_write_dataset("block-only-v1", [{
|
||||
"case_id": "c1",
|
||||
"query": "向量数据库",
|
||||
"expected_block_ids": [note.blocks[0].block_id],
|
||||
"citation_required": False,
|
||||
}])
|
||||
|
||||
run = _run(RAGRunRequest(dataset_id="block-only-v1", modes=[SearchMode.fts]))
|
||||
|
||||
assert run.status.value == "completed"
|
||||
fts = run.metrics["fts"]
|
||||
assert fts["hit_at_1"] == 1.0
|
||||
assert fts["recall_at_k"] == 1.0
|
||||
assert fts["mrr"] == 1.0
|
||||
|
||||
|
||||
def test_sse_stream_ends_on_terminal_event_in_replay() -> None:
|
||||
"""历史回放期间遇到终止事件时流应立即结束,而非进入实时队列永久等待(审阅 P2)。"""
|
||||
from app import routes
|
||||
from app.benchmarks import service
|
||||
from app.contracts import BenchmarkEvent, BenchmarkEventType
|
||||
|
||||
run_id = "benchmark_sse_replay"
|
||||
now = service._now()
|
||||
# 模拟「回放期间运行完成」:run 仍为 running(subscribe 返回非空队列),
|
||||
# 但历史事件里已含 RunCompleted 终止事件。
|
||||
service._runs[run_id] = BenchmarkRun(
|
||||
run_id=run_id,
|
||||
kind=BenchmarkKind.rag,
|
||||
dataset_id="d",
|
||||
dataset_hash="sha256:x",
|
||||
status=BenchmarkStatus.running,
|
||||
created_at=now,
|
||||
)
|
||||
service._events[run_id] = [
|
||||
BenchmarkEvent(
|
||||
event=BenchmarkEventType.run_started, run_id=run_id, sequence=0,
|
||||
data={}, timestamp=now,
|
||||
),
|
||||
BenchmarkEvent(
|
||||
event=BenchmarkEventType.run_completed, run_id=run_id, sequence=1,
|
||||
data={}, timestamp=now,
|
||||
),
|
||||
]
|
||||
try:
|
||||
# 直调路由函数时 FastAPI 不解析 Query/Header 默认值,需显式传 None 覆盖 Header 哨兵
|
||||
response = asyncio.run(
|
||||
routes.benchmark_events(run_id, after_sequence=-1, last_event_id=None)
|
||||
)
|
||||
|
||||
async def _collect() -> list[str]:
|
||||
out: list[str] = []
|
||||
async for chunk in response.body_iterator:
|
||||
out.append(chunk)
|
||||
return out
|
||||
|
||||
# 加超时防止回归(旧实现会永久挂起)
|
||||
chunks = asyncio.run(asyncio.wait_for(_collect(), timeout=5))
|
||||
finally:
|
||||
service._forget(run_id)
|
||||
|
||||
events = [
|
||||
line for chunk in chunks for line in chunk.splitlines() if line.startswith("event: ")
|
||||
]
|
||||
assert events == ["event: RunStarted", "event: RunCompleted"]
|
||||
@@ -0,0 +1,769 @@
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from app.agent.tools import ToolExecutionContext, ToolRegistry
|
||||
from app.config import BACKEND_DIR, get_settings
|
||||
from app.contracts import McpServerCreateRequest, McpServerUpdateRequest, ToolCall
|
||||
from app.extensions.mcp import McpLegacySseClient
|
||||
from app.extensions.mcp_registry import McpRegistryError, McpServerRegistry
|
||||
from app.providers.credentials import CredentialStoreError, EncryptedCredentialStore
|
||||
|
||||
SERVER = BACKEND_DIR / "extensions" / "fixtures" / "mcp-echo" / "server.py"
|
||||
|
||||
|
||||
def request(**overrides) -> McpServerCreateRequest:
|
||||
values = {
|
||||
"name": "Echo MCP",
|
||||
"command": sys.executable,
|
||||
"args": [str(SERVER)],
|
||||
"permissions": ["notes.read", "secrets.use"],
|
||||
"secret_environment_keys": ["TEST_MCP_SECRET"],
|
||||
}
|
||||
values.update(overrides)
|
||||
return McpServerCreateRequest(**values)
|
||||
|
||||
|
||||
def registry(*, launch: bool = True) -> McpServerRegistry:
|
||||
return McpServerRegistry(
|
||||
ToolRegistry(),
|
||||
EncryptedCredentialStore(),
|
||||
get_settings().data_dir,
|
||||
allow_process_launch=launch,
|
||||
)
|
||||
|
||||
|
||||
def test_registry_requires_current_trust_and_never_returns_secret() -> None:
|
||||
service = registry()
|
||||
created = service.create(request())
|
||||
assert created.trusted is False
|
||||
assert created.secret_environment == {"TEST_MCP_SECRET": False}
|
||||
|
||||
service.put_secret(created.server_id, "TEST_MCP_SECRET", "do-not-return")
|
||||
configured = service.get(created.server_id)
|
||||
assert configured.secret_environment == {"TEST_MCP_SECRET": True}
|
||||
assert "do-not-return" not in configured.model_dump_json()
|
||||
|
||||
with pytest.raises(McpRegistryError, match="approve"):
|
||||
service.test(created.server_id)
|
||||
|
||||
service.trust(created.server_id, created.command_digest)
|
||||
tested = service.test(created.server_id)
|
||||
assert tested.status == "stopped"
|
||||
assert tested.last_test_succeeded is True
|
||||
assert tested.tools_count > 0
|
||||
service.shutdown()
|
||||
|
||||
|
||||
def test_secret_change_disables_server_and_requires_a_new_connection_test() -> None:
|
||||
service = registry()
|
||||
created = service.create(request())
|
||||
service.put_secret(created.server_id, "TEST_MCP_SECRET", "first")
|
||||
service.trust(created.server_id, created.command_digest)
|
||||
service.test(created.server_id)
|
||||
service.enable(created.server_id)
|
||||
|
||||
service.put_secret(created.server_id, "TEST_MCP_SECRET", "second")
|
||||
current = service.get(created.server_id)
|
||||
assert current.enabled is False
|
||||
assert current.last_test_succeeded is None
|
||||
assert not any(
|
||||
item.name.startswith(f"mcp.{created.server_id}.")
|
||||
for item in service.tools.definitions()
|
||||
)
|
||||
with pytest.raises(McpRegistryError) as error:
|
||||
service.enable(created.server_id)
|
||||
assert error.value.code == "MCP_CONNECTION_TEST_REQUIRED"
|
||||
service.shutdown()
|
||||
|
||||
|
||||
def test_update_disables_server_and_revokes_command_trust() -> None:
|
||||
service = registry()
|
||||
created = service.create(request(secret_environment_keys=[]))
|
||||
service.trust(created.server_id, created.command_digest)
|
||||
service.test(created.server_id)
|
||||
enabled = service.enable(created.server_id)
|
||||
assert enabled.enabled is True
|
||||
assert any(
|
||||
item.name.startswith(f"mcp.{created.server_id}.")
|
||||
for item in service.tools.definitions()
|
||||
)
|
||||
|
||||
updated = service.update(
|
||||
created.server_id,
|
||||
McpServerUpdateRequest(
|
||||
**request(name="Changed", secret_environment_keys=[]).model_dump(),
|
||||
version=enabled.version,
|
||||
),
|
||||
)
|
||||
assert updated.enabled is False
|
||||
assert updated.trusted is False
|
||||
assert not any(
|
||||
item.name.startswith(f"mcp.{created.server_id}.")
|
||||
for item in service.tools.definitions()
|
||||
)
|
||||
service.shutdown()
|
||||
|
||||
|
||||
def test_update_remains_retryable_when_removed_secret_cleanup_fails(
|
||||
monkeypatch,
|
||||
) -> None:
|
||||
service = registry()
|
||||
created = service.create(request())
|
||||
service.put_secret(created.server_id, "TEST_MCP_SECRET", "keep-until-retry")
|
||||
|
||||
def fail_delete_many(_secret_ids: list[str]) -> set[str]:
|
||||
raise CredentialStoreError("credential store unavailable")
|
||||
|
||||
monkeypatch.setattr(service.credentials, "delete_many", fail_delete_many)
|
||||
with pytest.raises(McpRegistryError) as error:
|
||||
service.update(
|
||||
created.server_id,
|
||||
McpServerUpdateRequest(
|
||||
**request(secret_environment_keys=[]).model_dump(),
|
||||
version=created.version,
|
||||
),
|
||||
)
|
||||
|
||||
current = service.get(created.server_id)
|
||||
assert error.value.code == "MCP_SECRET_STORE_ERROR"
|
||||
assert current.version == created.version
|
||||
assert current.secret_environment == {"TEST_MCP_SECRET": True}
|
||||
service.shutdown()
|
||||
|
||||
|
||||
def test_delete_keeps_server_retryable_when_secret_cleanup_fails(monkeypatch) -> None:
|
||||
service = registry()
|
||||
created = service.create(request())
|
||||
service.put_secret(created.server_id, "TEST_MCP_SECRET", "keep-until-retry")
|
||||
|
||||
def fail_delete_many(_secret_ids: list[str]) -> set[str]:
|
||||
raise CredentialStoreError("credential store unavailable")
|
||||
|
||||
monkeypatch.setattr(service.credentials, "delete_many", fail_delete_many)
|
||||
with pytest.raises(McpRegistryError) as error:
|
||||
service.delete(created.server_id)
|
||||
|
||||
current = service.get(created.server_id)
|
||||
assert error.value.code == "MCP_SECRET_STORE_ERROR"
|
||||
assert current.server_id == created.server_id
|
||||
assert current.secret_environment == {"TEST_MCP_SECRET": True}
|
||||
service.shutdown()
|
||||
|
||||
|
||||
def test_unavailable_server_removes_bridge_host(monkeypatch) -> None:
|
||||
service = registry()
|
||||
created = service.create(request(secret_environment_keys=[]))
|
||||
with service._lock:
|
||||
service._records[created.server_id] = {
|
||||
**service._records[created.server_id],
|
||||
"enabled": True,
|
||||
}
|
||||
removed: list[str] = []
|
||||
monkeypatch.setattr(service.bridge, "remove", removed.append)
|
||||
|
||||
generation = object()
|
||||
service._generations[created.server_id] = generation
|
||||
service._unavailable(created.server_id, generation, "connection lost")
|
||||
|
||||
current = service.get(created.server_id)
|
||||
assert removed == [f"mcp.{created.server_id}"]
|
||||
assert current.enabled is False
|
||||
assert current.status == "unhealthy"
|
||||
service.shutdown()
|
||||
|
||||
|
||||
def test_old_failure_callback_cannot_stop_replacement_host(monkeypatch) -> None:
|
||||
service = registry()
|
||||
callbacks = []
|
||||
original_start = service.bridge.start
|
||||
|
||||
def capture_callback(*args, **kwargs):
|
||||
callbacks.append(args[4])
|
||||
return original_start(*args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(service.bridge, "start", capture_callback)
|
||||
created = service.create(request(secret_environment_keys=[]))
|
||||
service.trust(created.server_id, created.command_digest)
|
||||
callback_thread = None
|
||||
try:
|
||||
service.test(created.server_id)
|
||||
service.enable(created.server_id)
|
||||
old_callback = callbacks[-1]
|
||||
callback_started = threading.Event()
|
||||
callback_finished = threading.Event()
|
||||
|
||||
def delayed_failure():
|
||||
callback_started.set()
|
||||
old_callback(f"mcp.{created.server_id}", "delayed old failure")
|
||||
callback_finished.set()
|
||||
|
||||
# Queue the old callback while a replacement owns the lifecycle lock.
|
||||
with service._lifecycle_lock:
|
||||
callback_thread = threading.Thread(target=delayed_failure, daemon=True)
|
||||
callback_thread.start()
|
||||
assert callback_started.wait(timeout=2)
|
||||
service.disable(created.server_id)
|
||||
service.enable(created.server_id)
|
||||
assert callback_finished.wait(timeout=2)
|
||||
assert service.get(created.server_id).enabled is True
|
||||
assert service.get(created.server_id).status == "ready"
|
||||
assert service.tools.definitions()
|
||||
callbacks[-1](f"mcp.{created.server_id}", "current failure")
|
||||
assert service.get(created.server_id).enabled is False
|
||||
assert service.get(created.server_id).status == "unhealthy"
|
||||
finally:
|
||||
service.shutdown()
|
||||
if callback_thread is not None:
|
||||
callback_thread.join(timeout=2)
|
||||
|
||||
|
||||
def test_header_case_only_rename_preserves_secret() -> None:
|
||||
service = registry()
|
||||
config = {
|
||||
"name": "HTTP",
|
||||
"transport": "streamable_http",
|
||||
"url": "https://example.test/mcp",
|
||||
"secret_header_keys": ["Authorization"],
|
||||
}
|
||||
created = service.create(McpServerCreateRequest(**config))
|
||||
service.put_secret(created.server_id, "Authorization", "synthetic", kind="header")
|
||||
config["secret_header_keys"] = ["authorization"]
|
||||
updated = service.update(
|
||||
created.server_id, McpServerUpdateRequest(**config, version=created.version)
|
||||
)
|
||||
assert updated.secret_headers == {"authorization": True}
|
||||
assert (
|
||||
service.credentials.resolve(
|
||||
service._secret_id(created.server_id, "authorization", "header")
|
||||
)
|
||||
== "synthetic"
|
||||
)
|
||||
|
||||
|
||||
def test_environment_secrets_are_case_sensitive_and_delete_independently() -> None:
|
||||
service = registry()
|
||||
created = service.create(request(secret_environment_keys=["TOKEN", "token"]))
|
||||
service.put_secret(created.server_id, "TOKEN", "upper")
|
||||
service.put_secret(created.server_id, "token", "lower")
|
||||
assert (
|
||||
service.credentials.resolve(service._secret_id(created.server_id, "TOKEN"))
|
||||
== "upper"
|
||||
)
|
||||
assert (
|
||||
service.credentials.resolve(service._secret_id(created.server_id, "token"))
|
||||
== "lower"
|
||||
)
|
||||
service.delete_secret(created.server_id, "TOKEN")
|
||||
assert service.get(created.server_id).secret_environment == {
|
||||
"TOKEN": False,
|
||||
"token": True,
|
||||
}
|
||||
|
||||
|
||||
def test_legacy_environment_credential_migration_is_idempotent() -> None:
|
||||
service = registry()
|
||||
created = service.create(request(secret_environment_keys=["TOKEN"]))
|
||||
suffix = hashlib.sha256(b"environment\0token").hexdigest()[:20]
|
||||
legacy_id = f"mcp.{created.server_id}.{suffix}"
|
||||
service.credentials.put(legacy_id, "legacy-value")
|
||||
service._records[created.server_id]["secret_environment_version"] = 1
|
||||
service._write()
|
||||
migrated = registry()
|
||||
assert migrated.get(created.server_id).secret_environment == {"TOKEN": True}
|
||||
assert (
|
||||
migrated.credentials.resolve(migrated._secret_id(created.server_id, "TOKEN"))
|
||||
== "legacy-value"
|
||||
)
|
||||
assert not migrated.credentials.has(legacy_id)
|
||||
migrated.put_secret(created.server_id, "TOKEN", "new-value")
|
||||
assert (
|
||||
registry().credentials.resolve(migrated._secret_id(created.server_id, "TOKEN"))
|
||||
== "new-value"
|
||||
)
|
||||
|
||||
|
||||
def test_ambiguous_legacy_credentials_are_not_assigned_to_two_variables() -> None:
|
||||
service = registry()
|
||||
created = service.create(request(secret_environment_keys=["TOKEN", "token"]))
|
||||
suffix = hashlib.sha256(b"environment\0token").hexdigest()[:20]
|
||||
legacy_id = f"mcp.{created.server_id}.{suffix}"
|
||||
service.credentials.put(legacy_id, "cannot-reconstruct-originals")
|
||||
service._records[created.server_id]["secret_environment_version"] = 1
|
||||
service._write()
|
||||
migrated = registry()
|
||||
current = migrated.get(created.server_id)
|
||||
assert current.secret_environment == {"TOKEN": False, "token": False}
|
||||
assert current.enabled is False
|
||||
assert current.last_test_succeeded is None
|
||||
assert migrated.credentials.has(
|
||||
legacy_id
|
||||
) # Keep the original ciphertext recoverable.
|
||||
migrated.put_secret(created.server_id, "TOKEN", "upper")
|
||||
migrated.put_secret(created.server_id, "token", "lower")
|
||||
assert registry().get(created.server_id).secret_environment == {
|
||||
"TOKEN": True,
|
||||
"token": True,
|
||||
}
|
||||
migrated.delete(created.server_id)
|
||||
assert not migrated.credentials.has(legacy_id)
|
||||
|
||||
|
||||
def test_credential_id_migration_keeps_new_values_and_is_atomic(monkeypatch) -> None:
|
||||
credentials = EncryptedCredentialStore()
|
||||
credentials.put("mcp.old", "old-value")
|
||||
credentials.put("mcp.new", "new-value")
|
||||
original_write = credentials._write_tokens
|
||||
|
||||
def fail_write(_tokens):
|
||||
raise CredentialStoreError("synthetic failure")
|
||||
|
||||
monkeypatch.setattr(credentials, "_write_tokens", fail_write)
|
||||
with pytest.raises(CredentialStoreError):
|
||||
credentials.move_many({"mcp.old": "mcp.new"})
|
||||
assert credentials.resolve("mcp.old") == "old-value"
|
||||
assert credentials.resolve("mcp.new") == "new-value"
|
||||
monkeypatch.setattr(credentials, "_write_tokens", original_write)
|
||||
credentials.move_many({"mcp.old": "mcp.new"})
|
||||
assert credentials.resolve("mcp.old") is None
|
||||
assert credentials.resolve("mcp.new") == "new-value"
|
||||
|
||||
|
||||
def test_ambiguous_legacy_secret_is_not_resurrected_after_removing_a_key() -> None:
|
||||
service = registry()
|
||||
created = service.create(request(secret_environment_keys=["TOKEN", "token"]))
|
||||
legacy_id = service._legacy_environment_secret_id(created.server_id, "TOKEN")
|
||||
service.credentials.put(legacy_id, "ambiguous-old-value")
|
||||
service._records[created.server_id]["secret_environment_version"] = 1
|
||||
service._write()
|
||||
migrated = registry()
|
||||
migrated.update(
|
||||
created.server_id,
|
||||
McpServerUpdateRequest(
|
||||
**request(secret_environment_keys=["token"]).model_dump(),
|
||||
version=created.version,
|
||||
),
|
||||
)
|
||||
assert registry().get(created.server_id).secret_environment == {"token": False}
|
||||
|
||||
|
||||
def test_production_rejects_process_launch_even_after_approval() -> None:
|
||||
service = registry(launch=False)
|
||||
created = service.create(request(secret_environment_keys=[]))
|
||||
service.trust(created.server_id, created.command_digest)
|
||||
with pytest.raises(McpRegistryError) as error:
|
||||
service.enable(created.server_id)
|
||||
assert error.value.code == "MCP_SANDBOX_REQUIRED"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("startup,tool", [(120, 300), (1.5, 2.5)])
|
||||
def test_server_timeouts_survive_bridge_adaptation_and_reload(startup, tool) -> None:
|
||||
service = registry()
|
||||
created = service.create(
|
||||
request(
|
||||
secret_environment_keys=[],
|
||||
startup_timeout_seconds=startup,
|
||||
tool_timeout_seconds=tool,
|
||||
)
|
||||
)
|
||||
service.trust(created.server_id, created.command_digest)
|
||||
try:
|
||||
tested = service.test(created.server_id)
|
||||
assert tested.last_test_succeeded is True
|
||||
assert tested.startup_timeout_seconds == startup
|
||||
assert tested.tool_timeout_seconds == tool
|
||||
restored = registry().get(created.server_id)
|
||||
assert restored.startup_timeout_seconds == startup
|
||||
assert restored.tool_timeout_seconds == tool
|
||||
finally:
|
||||
service.shutdown()
|
||||
|
||||
|
||||
def test_enable_requires_successful_test_and_update_checks_version() -> None:
|
||||
service = registry()
|
||||
created = service.create(request(secret_environment_keys=[]))
|
||||
service.trust(created.server_id, created.command_digest)
|
||||
with pytest.raises(McpRegistryError) as error:
|
||||
service.enable(created.server_id)
|
||||
assert error.value.code == "MCP_CONNECTION_TEST_REQUIRED"
|
||||
|
||||
with pytest.raises(McpRegistryError) as error:
|
||||
service.update(
|
||||
created.server_id,
|
||||
McpServerUpdateRequest(
|
||||
**request(secret_environment_keys=[]).model_dump(), version=99
|
||||
),
|
||||
)
|
||||
assert error.value.code == "MCP_SERVER_VERSION_CONFLICT"
|
||||
|
||||
|
||||
def test_http_transport_rejects_invalid_cross_transport_fields() -> None:
|
||||
service = registry()
|
||||
with pytest.raises(McpRegistryError) as error:
|
||||
service.create(
|
||||
request(
|
||||
transport="streamable_http",
|
||||
url="https://example.invalid/mcp",
|
||||
secret_environment_keys=[],
|
||||
)
|
||||
)
|
||||
assert error.value.code == "MCP_CONFIG_INVALID"
|
||||
|
||||
|
||||
def test_registry_rejects_corrupt_persisted_json(tmp_path) -> None:
|
||||
path = tmp_path / "mcp"
|
||||
path.mkdir()
|
||||
(path / "servers.json").write_text("{broken", encoding="utf-8")
|
||||
with pytest.raises(McpRegistryError) as error:
|
||||
McpServerRegistry(
|
||||
ToolRegistry(),
|
||||
EncryptedCredentialStore(),
|
||||
tmp_path,
|
||||
allow_process_launch=True,
|
||||
)
|
||||
assert error.value.code == "MCP_REGISTRY_INVALID"
|
||||
|
||||
|
||||
def test_registry_rejects_structurally_invalid_record(tmp_path) -> None:
|
||||
path = tmp_path / "mcp"
|
||||
path.mkdir()
|
||||
(path / "servers.json").write_text(
|
||||
json.dumps({"server-1": {"name": "Broken", "transport": "stdio"}}),
|
||||
encoding="utf-8",
|
||||
)
|
||||
with pytest.raises(McpRegistryError) as error:
|
||||
McpServerRegistry(
|
||||
ToolRegistry(),
|
||||
EncryptedCredentialStore(),
|
||||
tmp_path,
|
||||
allow_process_launch=True,
|
||||
)
|
||||
assert error.value.code == "MCP_REGISTRY_INVALID"
|
||||
|
||||
|
||||
def test_registry_rejects_create_before_exceeding_persisted_limit(
|
||||
monkeypatch,
|
||||
) -> None:
|
||||
service = registry()
|
||||
service.create(request(name="Only server"))
|
||||
monkeypatch.setattr("app.extensions.mcp_registry._MAX_MCP_SERVERS", 1)
|
||||
with pytest.raises(McpRegistryError) as error:
|
||||
service.create(request(name="One too many"))
|
||||
assert error.value.code == "MCP_SERVER_LIMIT_REACHED"
|
||||
assert len(service.list()) == 1
|
||||
service.shutdown()
|
||||
|
||||
|
||||
def test_stdio_command_is_not_parsed_as_a_shell_string() -> None:
|
||||
service = registry()
|
||||
created = service.create(
|
||||
request(
|
||||
command=f'"{sys.executable}" "{SERVER}"',
|
||||
args=[],
|
||||
secret_environment_keys=[],
|
||||
)
|
||||
)
|
||||
service.trust(created.server_id, created.command_digest)
|
||||
with pytest.raises(McpRegistryError) as error:
|
||||
service.test(created.server_id)
|
||||
assert error.value.code == "PLUGIN_HOST_START_FAILED"
|
||||
assert service.get(created.server_id).last_test_succeeded is False
|
||||
service.shutdown()
|
||||
|
||||
|
||||
def test_enabled_server_is_restored_from_persisted_registry() -> None:
|
||||
first = registry()
|
||||
created = first.create(request(secret_environment_keys=[]))
|
||||
first.trust(created.server_id, created.command_digest)
|
||||
first.test(created.server_id)
|
||||
first.enable(created.server_id)
|
||||
first.shutdown()
|
||||
|
||||
restored = registry()
|
||||
restored.restore_enabled()
|
||||
current = restored.get(created.server_id)
|
||||
assert current.enabled is True
|
||||
assert current.status == "ready"
|
||||
assert any(
|
||||
item.name.startswith(f"mcp.{created.server_id}.")
|
||||
for item in restored.tools.definitions()
|
||||
)
|
||||
restored.shutdown()
|
||||
|
||||
|
||||
def test_lifecycle_operations_are_serialized_and_tool_names_are_isolated() -> None:
|
||||
service = registry()
|
||||
servers = [
|
||||
service.create(request(name=f"Echo {index}", secret_environment_keys=[]))
|
||||
for index in range(2)
|
||||
]
|
||||
for server in servers:
|
||||
service.trust(server.server_id, server.command_digest)
|
||||
service.test(server.server_id)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=4) as pool:
|
||||
enabled = list(
|
||||
pool.map(lambda item: service.enable(item.server_id), servers * 2)
|
||||
)
|
||||
assert all(item.enabled for item in enabled)
|
||||
names = [
|
||||
item.name for item in service.tools.definitions() if item.source == "mcp_server"
|
||||
]
|
||||
assert len(names) == len(set(names))
|
||||
assert all(
|
||||
any(name.startswith(f"mcp.{item.server_id}.") for name in names)
|
||||
for item in servers
|
||||
)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=4) as pool:
|
||||
list(pool.map(lambda item: service.disable(item.server_id), servers * 2))
|
||||
assert not any(item.source == "mcp_server" for item in service.tools.definitions())
|
||||
service.shutdown()
|
||||
|
||||
|
||||
def _http_result(request_id: int, result: dict) -> httpx.Response:
|
||||
return httpx.Response(
|
||||
200,
|
||||
headers={"content-type": "application/json"},
|
||||
json={"jsonrpc": "2.0", "id": request_id, "result": result},
|
||||
)
|
||||
|
||||
|
||||
def test_streamable_http_supports_session_headers_secrets_and_tool_summary(
|
||||
monkeypatch,
|
||||
) -> None:
|
||||
requests: list[httpx.Request] = []
|
||||
request_timeouts: dict[str, float] = {}
|
||||
|
||||
def handler(request_value: httpx.Request) -> httpx.Response:
|
||||
requests.append(request_value)
|
||||
if request_value.method == "GET":
|
||||
return httpx.Response(405)
|
||||
if request_value.method == "DELETE":
|
||||
return httpx.Response(405)
|
||||
payload = json.loads(request_value.content)
|
||||
timeout = request_value.extensions.get("timeout", {}).get("read")
|
||||
if isinstance(timeout, (int, float)):
|
||||
request_timeouts[payload.get("method", "notification")] = float(timeout)
|
||||
if payload.get("method") == "initialize":
|
||||
response = _http_result(
|
||||
payload["id"],
|
||||
{
|
||||
"protocolVersion": "2025-11-25",
|
||||
"capabilities": {"tools": {}},
|
||||
"serverInfo": {"name": "HTTP Fixture", "version": "1"},
|
||||
},
|
||||
)
|
||||
response.headers["MCP-Session-Id"] = "session-test"
|
||||
return response
|
||||
if payload.get("method") == "tools/list":
|
||||
return _http_result(
|
||||
payload["id"],
|
||||
{
|
||||
"tools": [
|
||||
{
|
||||
"name": "echo",
|
||||
"description": "Echo over HTTP",
|
||||
"inputSchema": {"type": "object", "properties": {}},
|
||||
}
|
||||
]
|
||||
},
|
||||
)
|
||||
if payload.get("method") == "tools/call":
|
||||
return _http_result(
|
||||
payload["id"], {"structuredContent": {"transport": "http"}}
|
||||
)
|
||||
return httpx.Response(202)
|
||||
|
||||
real_client = httpx.Client
|
||||
monkeypatch.setattr(
|
||||
"app.extensions.mcp.httpx.Client",
|
||||
lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs),
|
||||
)
|
||||
service = registry()
|
||||
created = service.create(
|
||||
McpServerCreateRequest(
|
||||
name="Remote MCP",
|
||||
transport="streamable_http",
|
||||
url="https://mcp.example.test/mcp",
|
||||
headers={"X-Client": "NotesAgent"},
|
||||
secret_header_keys=["Authorization"],
|
||||
)
|
||||
)
|
||||
service.put_secret(
|
||||
created.server_id, "Authorization", "Bearer hidden", kind="header"
|
||||
)
|
||||
service.trust(created.server_id, created.command_digest)
|
||||
tested = service.test(created.server_id)
|
||||
|
||||
assert tested.last_test_succeeded is True
|
||||
assert tested.secret_headers == {"Authorization": True}
|
||||
assert "Bearer hidden" not in tested.model_dump_json()
|
||||
assert service.list_tools(created.server_id)[0].remote_name == "echo"
|
||||
assert any(
|
||||
request.headers.get("mcp-session-id") == "session-test" for request in requests
|
||||
)
|
||||
assert any(
|
||||
request.headers.get("mcp-protocol-version") == "2025-11-25"
|
||||
for request in requests
|
||||
)
|
||||
assert all(
|
||||
request.headers.get("authorization") == "Bearer hidden" for request in requests
|
||||
)
|
||||
assert request_timeouts["initialize"] == 15
|
||||
assert request_timeouts["notifications/initialized"] == 15
|
||||
assert request_timeouts["tools/list"] == 15
|
||||
enabled = service.enable(created.server_id)
|
||||
tool_name = service.list_tools(created.server_id)[0].name
|
||||
result = asyncio.run(
|
||||
service.tools.execute(
|
||||
ToolCall(tool_call_id="call-1", name=tool_name, arguments={}),
|
||||
ToolExecutionContext(run_id="run-1"),
|
||||
)
|
||||
)
|
||||
assert enabled.enabled is True
|
||||
assert result.success is True
|
||||
assert result.output == {"transport": "http"}
|
||||
assert request_timeouts["tools/call"] == 30
|
||||
service.disable(created.server_id)
|
||||
service.shutdown()
|
||||
|
||||
|
||||
class _LegacyEventStream(httpx.SyncByteStream):
|
||||
def __init__(self) -> None:
|
||||
self.closed = threading.Event()
|
||||
|
||||
def __iter__(self):
|
||||
yield b"event: endpoint\ndata: /messages\n\n"
|
||||
time.sleep(0.1)
|
||||
initialize = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": 1,
|
||||
"result": {
|
||||
"protocolVersion": "2024-11-05",
|
||||
"capabilities": {"tools": {}},
|
||||
"serverInfo": {"name": "Legacy Fixture"},
|
||||
},
|
||||
}
|
||||
yield f"data: {json.dumps(initialize)}\n\n".encode()
|
||||
time.sleep(0.1)
|
||||
tools = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": 2,
|
||||
"result": {"tools": []},
|
||||
}
|
||||
yield f"data: {json.dumps(tools)}\n\n".encode()
|
||||
self.closed.wait()
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed.set()
|
||||
|
||||
|
||||
def test_legacy_sse_uses_same_origin_endpoint(monkeypatch) -> None:
|
||||
posted_urls: list[str] = []
|
||||
event_stream = _LegacyEventStream()
|
||||
|
||||
def handler(request_value: httpx.Request) -> httpx.Response:
|
||||
if request_value.method == "GET":
|
||||
return httpx.Response(
|
||||
200,
|
||||
headers={"content-type": "text/event-stream"},
|
||||
stream=event_stream,
|
||||
)
|
||||
posted_urls.append(str(request_value.url))
|
||||
return httpx.Response(202)
|
||||
|
||||
real_client = httpx.Client
|
||||
monkeypatch.setattr(
|
||||
"app.extensions.mcp.httpx.Client",
|
||||
lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs),
|
||||
)
|
||||
service = registry()
|
||||
created = service.create(
|
||||
McpServerCreateRequest(
|
||||
name="Legacy MCP",
|
||||
transport="sse",
|
||||
url="https://legacy.example.test/sse",
|
||||
)
|
||||
)
|
||||
service.trust(created.server_id, created.command_digest)
|
||||
tested = service.test(created.server_id)
|
||||
assert tested.last_test_succeeded is True
|
||||
assert posted_urls and all(
|
||||
url == "https://legacy.example.test/messages" for url in posted_urls
|
||||
)
|
||||
service.shutdown()
|
||||
event_stream.close()
|
||||
|
||||
|
||||
class _EndingLegacyEventStream(httpx.SyncByteStream):
|
||||
def __iter__(self):
|
||||
yield b"event: endpoint\ndata: /messages\n\n"
|
||||
|
||||
|
||||
def test_legacy_sse_eof_marks_client_unavailable(monkeypatch) -> None:
|
||||
def handler(_request_value: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(
|
||||
200,
|
||||
headers={"content-type": "text/event-stream"},
|
||||
stream=_EndingLegacyEventStream(),
|
||||
)
|
||||
|
||||
real_client = httpx.Client
|
||||
monkeypatch.setattr(
|
||||
"app.extensions.mcp.httpx.Client",
|
||||
lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs),
|
||||
)
|
||||
broken = threading.Event()
|
||||
client = McpLegacySseClient(
|
||||
"https://legacy.example.test/sse",
|
||||
headers={},
|
||||
startup_timeout_seconds=1,
|
||||
on_seen=lambda: None,
|
||||
on_broken=lambda _message: broken.set(),
|
||||
on_tools_changed=lambda: None,
|
||||
)
|
||||
client.start()
|
||||
assert broken.wait(timeout=1)
|
||||
client.stop()
|
||||
|
||||
|
||||
class _CrossOriginLegacyEventStream(httpx.SyncByteStream):
|
||||
def __iter__(self):
|
||||
yield b"event: endpoint\ndata: https://attacker.example/messages\n\n"
|
||||
|
||||
|
||||
def test_legacy_sse_rejects_cross_origin_message_endpoint(monkeypatch) -> None:
|
||||
def handler(request_value: httpx.Request) -> httpx.Response:
|
||||
assert request_value.method == "GET"
|
||||
return httpx.Response(
|
||||
200,
|
||||
headers={"content-type": "text/event-stream"},
|
||||
stream=_CrossOriginLegacyEventStream(),
|
||||
)
|
||||
|
||||
real_client = httpx.Client
|
||||
monkeypatch.setattr(
|
||||
"app.extensions.mcp.httpx.Client",
|
||||
lambda **kwargs: real_client(transport=httpx.MockTransport(handler), **kwargs),
|
||||
)
|
||||
service = registry()
|
||||
created = service.create(
|
||||
McpServerCreateRequest(
|
||||
name="Unsafe legacy MCP",
|
||||
transport="sse",
|
||||
url="https://legacy.example.test/sse",
|
||||
)
|
||||
)
|
||||
service.trust(created.server_id, created.command_digest)
|
||||
with pytest.raises(McpRegistryError) as error:
|
||||
service.test(created.server_id)
|
||||
assert error.value.code == "MCP_HTTP_RESPONSE_INVALID"
|
||||
service.shutdown()
|
||||
@@ -0,0 +1,83 @@
|
||||
import json
|
||||
from contextlib import closing
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from app.extensions import mcp
|
||||
|
||||
|
||||
class ChunkStream(httpx.SyncByteStream):
|
||||
def __init__(self, chunks):
|
||||
self.chunks = chunks
|
||||
self.bytes_read = 0
|
||||
|
||||
def __iter__(self):
|
||||
for chunk in self.chunks:
|
||||
self.bytes_read += len(chunk)
|
||||
yield chunk
|
||||
|
||||
|
||||
def test_sse_rejects_unterminated_line_before_reading_entire_stream(monkeypatch):
|
||||
monkeypatch.setattr(mcp, "MAX_MCP_MESSAGE_BYTES", 1024)
|
||||
stream = ChunkStream([b"x" * 256] * 256)
|
||||
with (
|
||||
closing(httpx.Response(200, stream=stream)) as response,
|
||||
pytest.raises(mcp.McpBridgeError, match="too large"),
|
||||
):
|
||||
list(mcp._iter_sse(response))
|
||||
assert stream.bytes_read == 1280
|
||||
|
||||
|
||||
def test_sse_limits_combined_event_before_partial_line_is_complete(monkeypatch):
|
||||
monkeypatch.setattr(mcp, "MAX_MCP_MESSAGE_BYTES", 32)
|
||||
stream = ChunkStream(
|
||||
[b"data: 123456789\n", b"data: 123456789\n", b"x", b"not-read"]
|
||||
)
|
||||
with (
|
||||
closing(httpx.Response(200, stream=stream)) as response,
|
||||
pytest.raises(mcp.McpBridgeError, match="too large"),
|
||||
):
|
||||
list(mcp._iter_sse(response))
|
||||
assert stream.bytes_read == 33
|
||||
|
||||
|
||||
@pytest.mark.parametrize("separator", [b"\n", b"\r", b"\r\n"])
|
||||
@pytest.mark.parametrize("chunk_size", [1, 2, 7, 1024])
|
||||
def test_sse_preserves_utf8_and_line_endings_across_chunks(separator, chunk_size):
|
||||
payload = json.dumps(
|
||||
{"jsonrpc": "2.0", "id": 1, "result": {"text": "中文"}}, ensure_ascii=False
|
||||
)
|
||||
wire = b"\xef\xbb\xbf" + separator.join(
|
||||
[
|
||||
b": heartbeat",
|
||||
b"event: message",
|
||||
b"id: replay-1",
|
||||
("data: " + payload).encode(),
|
||||
b"",
|
||||
b"",
|
||||
]
|
||||
)
|
||||
stream = ChunkStream(
|
||||
[wire[index : index + chunk_size] for index in range(0, len(wire), chunk_size)]
|
||||
)
|
||||
with closing(httpx.Response(200, stream=stream)) as response:
|
||||
assert list(mcp._iter_sse(response)) == [("message", "replay-1", payload)]
|
||||
|
||||
|
||||
def test_sse_event_limit_resets_between_events(monkeypatch):
|
||||
monkeypatch.setattr(mcp, "MAX_MCP_MESSAGE_BYTES", 16)
|
||||
with closing(
|
||||
httpx.Response(200, stream=ChunkStream([b"data: one\n\ndata: two\r\r"]))
|
||||
) as response:
|
||||
assert list(mcp._iter_sse(response)) == [
|
||||
("message", None, "one"),
|
||||
("message", None, "two"),
|
||||
]
|
||||
|
||||
|
||||
def test_sse_preserves_multiline_data_and_final_unterminated_line():
|
||||
with closing(
|
||||
httpx.Response(200, stream=ChunkStream([b"data: first\ndata: last"]))
|
||||
) as response:
|
||||
assert list(mcp._iter_sse(response)) == [("message", None, "first\nlast")]
|
||||
@@ -0,0 +1,663 @@
|
||||
"""Offline model-routing contracts, HTTP validation, media lifetimes and persistence.
|
||||
|
||||
All HTTP uses MockTransport (or the in-process API). Credentials, models and
|
||||
attachments are fakes, and conftest redirects all storage to temporary paths.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
from email import policy
|
||||
from email.parser import BytesParser
|
||||
from types import SimpleNamespace
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from app.contracts import ModelBinding, ModelRoutingConfig, ProviderConfig, ProviderType
|
||||
from app.errors import ApiError
|
||||
from app.providers import MockProvider
|
||||
from app.providers.credentials import CredentialStoreError
|
||||
from app.providers.registry import ProviderRegistry
|
||||
from app.providers.routing import ModelRoutingService, PendingSpeechBackend
|
||||
from app.retrieval.embedding import HashEmbeddingProvider
|
||||
|
||||
|
||||
def run(awaitable):
|
||||
return asyncio.run(awaitable)
|
||||
|
||||
|
||||
def response(data, status=200):
|
||||
# Raw JSON intentionally permits NaN/Infinity to exercise hostile API output.
|
||||
return httpx.Response(status, content=json.dumps(data).encode(), headers={"content-type": "application/json"})
|
||||
|
||||
|
||||
class FakeCredentials:
|
||||
def __init__(self):
|
||||
self.value = "unit-test-placeholder"
|
||||
self.error = None
|
||||
self.calls = []
|
||||
|
||||
def resolve(self, credential_id):
|
||||
self.calls.append(credential_id)
|
||||
if self.error:
|
||||
raise self.error
|
||||
return self.value if credential_id else None
|
||||
|
||||
|
||||
class FakeEmbedding:
|
||||
model_id = "fake-local-model"
|
||||
dim = 3
|
||||
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
self.error = None
|
||||
|
||||
async def embed_documents(self, texts):
|
||||
self.calls.append(list(texts))
|
||||
if self.error:
|
||||
raise self.error
|
||||
return [[0.6, 0.8, 0.0] for _ in texts]
|
||||
|
||||
|
||||
class FakeSpeech:
|
||||
available = True
|
||||
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
self.text = "local transcript"
|
||||
self.score = 0.25
|
||||
self.error = None
|
||||
|
||||
async def transcribe(self, source, language):
|
||||
self.calls.append(("transcribe", source, language))
|
||||
if self.error:
|
||||
raise self.error
|
||||
return self.text
|
||||
|
||||
async def match(self, source, reference):
|
||||
self.calls.append(("match", source, reference))
|
||||
if self.error:
|
||||
raise self.error
|
||||
return self.score
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def no_real_http(monkeypatch):
|
||||
async def reject_async(*args, **kwargs):
|
||||
pytest.fail("Real HTTP transport is forbidden in model-routing tests")
|
||||
|
||||
def reject_sync(*args, **kwargs):
|
||||
pytest.fail("Real HTTP transport is forbidden in model-routing tests")
|
||||
|
||||
monkeypatch.setattr(httpx.AsyncHTTPTransport, "handle_async_request", reject_async)
|
||||
monkeypatch.setattr(httpx.HTTPTransport, "handle_request", reject_sync)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def rig():
|
||||
requests = []
|
||||
|
||||
def unexpected(request):
|
||||
pytest.fail(f"Unexpected model HTTP request: {request.url}")
|
||||
|
||||
state = SimpleNamespace(handler=unexpected)
|
||||
|
||||
async def dispatch(request):
|
||||
requests.append(request)
|
||||
result = state.handler(request)
|
||||
return await result if hasattr(result, "__await__") else result
|
||||
|
||||
providers = ProviderRegistry()
|
||||
config = ProviderConfig(
|
||||
provider_id="test-provider", provider_type=ProviderType.openai_compatible,
|
||||
name="Fake provider", base_url="https://models.invalid/v1/", credential_id="test-credential",
|
||||
)
|
||||
providers.register(config, MockProvider())
|
||||
credentials, embedding, speech = FakeCredentials(), FakeEmbedding(), FakeSpeech()
|
||||
service = ModelRoutingService(
|
||||
providers, credentials, local_embedding=embedding, local_speech=speech,
|
||||
transport=httpx.MockTransport(dispatch),
|
||||
)
|
||||
return SimpleNamespace(
|
||||
service=service, providers=providers, credentials=credentials,
|
||||
embedding=embedding, speech=speech, requests=requests, http=state,
|
||||
)
|
||||
|
||||
|
||||
def bind(rig, capability="embedding", **overrides):
|
||||
endpoints = {
|
||||
"embedding": "/embeddings", "transcription": "/audio/transcriptions",
|
||||
"speaker_matching": "/audio/speaker-matches",
|
||||
}
|
||||
binding = ModelBinding(**{
|
||||
"provider_id": "test-provider", "model": "test-model",
|
||||
"endpoint": endpoints[capability], **overrides,
|
||||
})
|
||||
current = rig.service.configuration()
|
||||
return rig.service.update(current.model_copy(update={capability: binding}))
|
||||
|
||||
|
||||
def assert_local(rig, result, texts, reason):
|
||||
assert result.source == "local"
|
||||
assert result.model_id == rig.embedding.model_id
|
||||
assert result.dimensions == 3
|
||||
assert result.vectors == [[0.6, 0.8, 0.0] for _ in texts]
|
||||
assert result.fallback_reason == reason
|
||||
assert rig.embedding.calls == [texts]
|
||||
|
||||
|
||||
def test_embedding_observation_keeps_request_binding_when_config_changes(rig):
|
||||
from app.retrieval.provenance import capture_embedding
|
||||
initial = bind(rig, model="original-model")
|
||||
|
||||
def handler(request):
|
||||
assert json.loads(request.content)["model"] == "original-model"
|
||||
bind(rig, model="next-model")
|
||||
return response({"data": [{"index": 0, "embedding": [1, 0, 0]}]})
|
||||
|
||||
rig.http.handler = handler
|
||||
with capture_embedding() as observation:
|
||||
result = run(rig.service.embed(["query"]))
|
||||
assert result.source == "api"
|
||||
assert observation["route_version"] == initial.config.version
|
||||
assert observation["requested_route"]["model"] == "original-model"
|
||||
assert observation["requested_route"]["provider_id"] == "test-provider"
|
||||
assert rig.service.configuration().embedding.model == "next-model"
|
||||
assert rig.credentials.value not in json.dumps(observation)
|
||||
assert "credential_id" not in json.dumps(observation)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def audio(tmp_path):
|
||||
source, reference = tmp_path / "audio.wav", tmp_path / "reference.wav"
|
||||
source.write_bytes(b"fake-audio-content")
|
||||
reference.write_bytes(b"fake-reference-content")
|
||||
return source, reference
|
||||
|
||||
|
||||
def media_call(rig, capability, audio):
|
||||
if capability == "transcription":
|
||||
return rig.service.transcribe(audio[0], "zh")
|
||||
return rig.service.match_speakers(*audio)
|
||||
|
||||
|
||||
def track_media_handles(rig, monkeypatch):
|
||||
handles = []
|
||||
original = rig.service._media_file
|
||||
|
||||
def tracked(path):
|
||||
handle = original(path)
|
||||
handles.append(handle)
|
||||
return handle
|
||||
|
||||
monkeypatch.setattr(rig.service, "_media_file", tracked)
|
||||
return handles
|
||||
|
||||
|
||||
def test_absent_binding_uses_hash_without_network(rig):
|
||||
rig.service.local_embedding = HashEmbeddingProvider()
|
||||
texts = ["hello retrieval", "向量检索"]
|
||||
result = run(rig.service.embed(texts))
|
||||
assert result.source == "local"
|
||||
assert result.model_id == "hash-v1"
|
||||
assert result.dimensions == 128
|
||||
assert result.vectors == run(HashEmbeddingProvider().embed_documents(texts))
|
||||
assert result.fallback_reason is None
|
||||
assert rig.requests == rig.credentials.calls == []
|
||||
statuses = {item.capability: item.status for item in rig.service.describe().local_backends}
|
||||
assert statuses == {"embedding": "placeholder", "transcription": "ready", "speaker_matching": "ready"}
|
||||
|
||||
|
||||
def test_empty_embedding_input_does_not_call_remote(rig):
|
||||
bind(rig)
|
||||
result = run(rig.service.embed([]))
|
||||
assert result.vectors == [] and result.source == "local"
|
||||
assert rig.requests == []
|
||||
|
||||
|
||||
def test_remote_embedding_restores_batch_order_normalizes_and_sends_auth(rig):
|
||||
bind(rig, dimensions=2)
|
||||
texts = [str(index) for index in range(35)]
|
||||
|
||||
def handler(request):
|
||||
assert request.method == "POST"
|
||||
assert str(request.url) == "https://models.invalid/v1/embeddings"
|
||||
assert request.headers["authorization"] == "Bearer unit-test-placeholder"
|
||||
payload = json.loads(request.content)
|
||||
assert payload["model"] == "test-model"
|
||||
assert payload["dimensions"] == 2
|
||||
assert payload["encoding_format"] == "float"
|
||||
return response({"data": [
|
||||
{"index": index, "embedding": [float(int(text) + 1), 1.0]}
|
||||
for index, text in reversed(list(enumerate(payload["input"])))
|
||||
]})
|
||||
|
||||
rig.http.handler = handler
|
||||
result = run(rig.service.embed(texts))
|
||||
assert result.source == "api" and result.fallback_reason is None
|
||||
assert result.dimensions == 2 and len(result.vectors) == 35
|
||||
for index, vector in enumerate(result.vectors):
|
||||
assert sum(value * value for value in vector) == pytest.approx(1.0)
|
||||
assert vector[0] / vector[1] == pytest.approx(index + 1)
|
||||
assert [json.loads(req.content)["input"] for req in rig.requests] == [texts[:32], texts[32:]]
|
||||
assert rig.embedding.calls == []
|
||||
|
||||
|
||||
def test_space_id_is_stable_and_includes_full_url_model_and_inferred_dimensions(rig):
|
||||
dimensions = 2
|
||||
|
||||
def handler(request):
|
||||
assert "dimensions" not in json.loads(request.content)
|
||||
return response({"data": [{"index": 0, "embedding": [1.0] * dimensions}]})
|
||||
|
||||
rig.http.handler = handler
|
||||
bind(rig, model=" trimmed-model ")
|
||||
|
||||
def check(url, model, dimension):
|
||||
result = run(rig.service.embed(["hello"]))
|
||||
digest = hashlib.sha256(json.dumps([url, model, dimension], separators=(",", ":")).encode()).hexdigest()
|
||||
assert result.model_id == "api-" + digest
|
||||
assert result.source == "api"
|
||||
return result.model_id
|
||||
|
||||
first = check("https://models.invalid/v1/embeddings", "trimmed-model", 2)
|
||||
assert first == check("https://models.invalid/v1/embeddings", "trimmed-model", 2)
|
||||
config = rig.providers.get_any("test-provider").config.model_copy(update={"base_url": "https://models.invalid/v1"})
|
||||
rig.providers.replace(config, MockProvider())
|
||||
assert first == check("https://models.invalid/v1/embeddings", "trimmed-model", 2)
|
||||
bind(rig, model="trimmed-model", endpoint="/custom/embeddings")
|
||||
endpoint_id = check("https://models.invalid/v1/custom/embeddings", "trimmed-model", 2)
|
||||
bind(rig, model="another-model", endpoint="/custom/embeddings")
|
||||
model_id = check("https://models.invalid/v1/custom/embeddings", "another-model", 2)
|
||||
dimensions = 3
|
||||
dim_id = check("https://models.invalid/v1/custom/embeddings", "another-model", 3)
|
||||
config = config.model_copy(update={"base_url": "https://other.invalid/v1"})
|
||||
rig.providers.replace(config, MockProvider())
|
||||
provider_id = check("https://other.invalid/v1/custom/embeddings", "another-model", 3)
|
||||
assert len({first, endpoint_id, model_id, dim_id, provider_id}) == 5
|
||||
|
||||
|
||||
@pytest.mark.parametrize("data", [
|
||||
{"data": []},
|
||||
{"data": [{"index": 0, "embedding": [1, 0]}]},
|
||||
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 0, "embedding": [0, 1]}]},
|
||||
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 2, "embedding": [0, 1]}]},
|
||||
{"data": [{"index": False, "embedding": [1, 0]}, {"index": 1, "embedding": [0, 1]}]},
|
||||
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 1, "embedding": [0, 1, 0]}]},
|
||||
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 1, "embedding": [float("nan"), 1]}]},
|
||||
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 1, "embedding": [float("inf"), 1]}]},
|
||||
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 1, "embedding": [True, 1]}]},
|
||||
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 1, "embedding": [0, 0]}]},
|
||||
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 1, "embedding": []}]},
|
||||
{"data": [{"index": 0, "embedding": [1, 0]}, {"index": 1, "embedding": ["1", 0]}]},
|
||||
{"data": [None, None]},
|
||||
{"error": {"message": "in-band failure"}, "data": []},
|
||||
[],
|
||||
], ids=["empty", "count", "duplicate-index", "out-of-range-index", "bool-index", "dimensions", "nan", "infinity", "bool", "zero", "empty-vector", "string", "invalid-items", "in-band-error", "non-object"])
|
||||
def test_invalid_remote_embeddings_fall_back_as_a_whole(rig, data):
|
||||
bind(rig)
|
||||
rig.http.handler = lambda request: response(data)
|
||||
texts = ["first", "second"]
|
||||
assert_local(rig, run(rig.service.embed(texts)), texts, "PROVIDER_INVALID_RESPONSE")
|
||||
|
||||
|
||||
def test_explicit_embedding_dimension_mismatch_falls_back(rig):
|
||||
bind(rig, dimensions=3)
|
||||
rig.http.handler = lambda request: response({"data": [{"index": 0, "embedding": [1, 0]}]})
|
||||
assert_local(rig, run(rig.service.embed(["text"])), ["text"], "PROVIDER_INVALID_RESPONSE")
|
||||
|
||||
|
||||
def test_later_batch_dimension_mismatch_discards_earlier_remote_vectors(rig):
|
||||
bind(rig)
|
||||
|
||||
def handler(request):
|
||||
batch = json.loads(request.content)["input"]
|
||||
dimension = 2 if len(rig.requests) == 1 else 3
|
||||
return response({"data": [{"index": i, "embedding": [1] * dimension} for i in range(len(batch))]})
|
||||
|
||||
rig.http.handler = handler
|
||||
texts = [str(i) for i in range(33)]
|
||||
assert_local(rig, run(rig.service.embed(texts)), texts, "PROVIDER_INVALID_RESPONSE")
|
||||
assert len(rig.requests) == 2
|
||||
|
||||
|
||||
@pytest.mark.parametrize("failure, reason", [
|
||||
(401, "PROVIDER_AUTH_FAILED"), (403, "PROVIDER_AUTH_FAILED"),
|
||||
(404, "MODEL_NOT_FOUND"), (429, "PROVIDER_RATE_LIMITED"), (500, "PROVIDER_UNAVAILABLE"),
|
||||
("timeout", "PROVIDER_TIMEOUT"), ("connect", "PROVIDER_UNAVAILABLE"),
|
||||
("json", "PROVIDER_INVALID_RESPONSE"),
|
||||
])
|
||||
def test_embedding_http_failures_use_injected_local(rig, failure, reason):
|
||||
bind(rig)
|
||||
|
||||
def handler(request):
|
||||
if failure == "timeout":
|
||||
raise httpx.ReadTimeout("simulated timeout", request=request)
|
||||
if failure == "connect":
|
||||
raise httpx.ConnectError("simulated connection failure", request=request)
|
||||
if failure == "json":
|
||||
return httpx.Response(200, content=b"not JSON")
|
||||
return response({"error": "failed"}, failure)
|
||||
|
||||
rig.http.handler = handler
|
||||
assert_local(rig, run(rig.service.embed(["text"])), ["text"], reason)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("failure, reason", [
|
||||
("missing-key", "PROVIDER_CREDENTIAL_MISSING"),
|
||||
("unreadable-key", "PROVIDER_CREDENTIAL_UNAVAILABLE"),
|
||||
("disabled-provider", "PROVIDER_UNAVAILABLE"),
|
||||
])
|
||||
def test_unavailable_remote_configuration_falls_back_without_http(rig, failure, reason):
|
||||
bind(rig)
|
||||
if failure == "missing-key":
|
||||
rig.credentials.value = None
|
||||
elif failure == "unreadable-key":
|
||||
rig.credentials.error = CredentialStoreError("fake unavailable store")
|
||||
else:
|
||||
config = rig.providers.get_any("test-provider").config.model_copy(update={"enabled": False})
|
||||
rig.providers.replace(config, MockProvider())
|
||||
assert_local(rig, run(rig.service.embed(["text"])), ["text"], reason)
|
||||
assert rig.requests == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize("capability", ["transcription", "speaker_matching"])
|
||||
def test_media_success_sends_expected_multipart_and_closes_files(rig, audio, monkeypatch, capability):
|
||||
bind(rig, capability)
|
||||
handles = track_media_handles(rig, monkeypatch)
|
||||
|
||||
def handler(request):
|
||||
assert request.headers["authorization"] == "Bearer unit-test-placeholder"
|
||||
assert str(request.url).endswith("/audio/transcriptions" if capability == "transcription" else "/audio/speaker-matches")
|
||||
message = BytesParser(policy=policy.default).parsebytes(
|
||||
b"Content-Type: " + request.headers["content-type"].encode() + b"\r\nMIME-Version: 1.0\r\n\r\n" + request.content,
|
||||
)
|
||||
parts = {part.get_param("name", header="content-disposition"): part for part in message.iter_parts()}
|
||||
assert parts["model"].get_payload(decode=True) == b"test-model"
|
||||
assert parts["file"].get_filename() == audio[0].name
|
||||
assert parts["file"].get_payload(decode=True) == audio[0].read_bytes()
|
||||
if capability == "transcription":
|
||||
assert set(parts) == {"model", "language", "file"}
|
||||
assert parts["language"].get_payload(decode=True) == b"zh"
|
||||
return response({"text": "remote transcript"})
|
||||
assert set(parts) == {"model", "file", "reference_file"}
|
||||
assert parts["reference_file"].get_filename() == audio[1].name
|
||||
assert parts["reference_file"].get_payload(decode=True) == audio[1].read_bytes()
|
||||
return response({"score": 0.875})
|
||||
|
||||
rig.http.handler = handler
|
||||
result = run(media_call(rig, capability, audio))
|
||||
assert result.source == "api" and result.fallback_reason is None
|
||||
assert result.text == "remote transcript" if capability == "transcription" else result.score == 0.875
|
||||
assert len(handles) == (1 if capability == "transcription" else 2)
|
||||
assert all(handle.closed for handle in handles)
|
||||
assert rig.speech.calls == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize("capability, data", [
|
||||
("transcription", {}), ("transcription", {"text": " "}), ("transcription", {"text": False}),
|
||||
("transcription", {"error": "in-band", "text": "must not use"}),
|
||||
("speaker_matching", {}), ("speaker_matching", {"score": -0.1}),
|
||||
("speaker_matching", {"score": 1.1}), ("speaker_matching", {"score": True}),
|
||||
("speaker_matching", {"score": float("nan")}), ("speaker_matching", {"score": "0.5"}),
|
||||
("speaker_matching", {"error": "in-band", "score": 0.9}),
|
||||
])
|
||||
def test_invalid_remote_media_falls_back_to_injected_local(rig, audio, monkeypatch, capability, data):
|
||||
bind(rig, capability)
|
||||
handles = track_media_handles(rig, monkeypatch)
|
||||
rig.http.handler = lambda request: response(data)
|
||||
result = run(media_call(rig, capability, audio))
|
||||
assert result.source == "local" and result.fallback_reason == "PROVIDER_INVALID_RESPONSE"
|
||||
assert result.text == "local transcript" if capability == "transcription" else result.score == 0.25
|
||||
assert rig.speech.calls == [
|
||||
("transcribe", audio[0], "zh") if capability == "transcription" else ("match", *audio)
|
||||
]
|
||||
assert handles and all(handle.closed for handle in handles)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("capability", ["transcription", "speaker_matching"])
|
||||
@pytest.mark.parametrize("configured", [False, True])
|
||||
def test_pending_local_backend_has_explicit_503_and_fallback_details(rig, audio, capability, configured):
|
||||
rig.service.local_speech = PendingSpeechBackend()
|
||||
if configured:
|
||||
bind(rig, capability)
|
||||
rig.http.handler = lambda request: response({"error": "unauthorized"}, 401)
|
||||
with pytest.raises(ApiError) as caught:
|
||||
run(media_call(rig, capability, audio))
|
||||
assert caught.value.status_code == 503
|
||||
assert caught.value.code == "LOCAL_MODEL_NOT_INSTALLED"
|
||||
assert caught.value.details == {"fallback_reason": "PROVIDER_AUTH_FAILED" if configured else None}
|
||||
statuses = {item.capability: item.status for item in rig.service.describe().local_backends}
|
||||
assert statuses["transcription"] == statuses["speaker_matching"] == "not_installed"
|
||||
assert len(rig.requests) == int(configured)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("capability", ["transcription", "speaker_matching"])
|
||||
def test_invalid_local_speech_returns_explicit_503(rig, audio, capability):
|
||||
rig.speech.text = ""
|
||||
rig.speech.score = True
|
||||
with pytest.raises(ApiError) as caught:
|
||||
run(media_call(rig, capability, audio))
|
||||
assert (caught.value.status_code, caught.value.code) == (503, "LOCAL_MODEL_INVALID_RESPONSE")
|
||||
assert caught.value.details == {"fallback_reason": None}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("capability", ["embedding", "transcription", "speaker_matching"])
|
||||
@pytest.mark.parametrize("stage", ["remote", "local"])
|
||||
def test_cancellation_propagates_and_upload_handles_close(rig, audio, monkeypatch, capability, stage):
|
||||
bind(rig, capability)
|
||||
handles = track_media_handles(rig, monkeypatch)
|
||||
|
||||
async def cancelled(request):
|
||||
raise asyncio.CancelledError()
|
||||
|
||||
if stage == "remote":
|
||||
rig.http.handler = cancelled
|
||||
else:
|
||||
rig.http.handler = lambda request: response({"error": "fallback"}, 500)
|
||||
rig.embedding.error = rig.speech.error = asyncio.CancelledError()
|
||||
operation = rig.service.embed(["text"]) if capability == "embedding" else media_call(rig, capability, audio)
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
run(operation)
|
||||
assert len(handles) == {"embedding": 0, "transcription": 1, "speaker_matching": 2}[capability]
|
||||
assert all(handle.closed for handle in handles)
|
||||
if stage == "remote":
|
||||
assert rig.embedding.calls == rig.speech.calls == []
|
||||
|
||||
|
||||
def test_missing_reference_closes_already_open_source(rig, audio, monkeypatch):
|
||||
bind(rig, "speaker_matching")
|
||||
handles = track_media_handles(rig, monkeypatch)
|
||||
audio[1].unlink()
|
||||
with pytest.raises(ApiError) as caught:
|
||||
run(rig.service.match_speakers(*audio))
|
||||
assert caught.value.status_code == 404
|
||||
assert len(handles) == 1 and handles[0].closed
|
||||
assert rig.requests == []
|
||||
|
||||
|
||||
def test_config_optimistic_conflict_preserves_saved_bindings(rig):
|
||||
assert rig.service.configuration().version == 0
|
||||
saved = bind(rig).config
|
||||
assert saved.version == 1
|
||||
with pytest.raises(ApiError) as caught:
|
||||
rig.service.update(ModelRoutingConfig(version=0))
|
||||
assert (caught.value.status_code, caught.value.code) == (409, "MODEL_ROUTING_VERSION_CONFLICT")
|
||||
assert rig.service.configuration() == saved
|
||||
assert rig.service.uses_provider("test-provider")
|
||||
assert not rig.service.uses_provider("not-a-provider")
|
||||
cleared = rig.service.update(ModelRoutingConfig(version=1)).config
|
||||
assert cleared.version == 2 and cleared.embedding is None
|
||||
assert not rig.service.uses_provider("test-provider")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("capability", ["embedding", "transcription", "speaker_matching"])
|
||||
@pytest.mark.parametrize("provider_id, code", [
|
||||
("missing", "PROVIDER_NOT_FOUND"), ("unsupported", "MODEL_ROUTING_PROTOCOL_UNSUPPORTED"),
|
||||
])
|
||||
def test_config_references_require_existing_supported_providers(rig, capability, provider_id, code):
|
||||
rig.providers.register(
|
||||
ProviderConfig(provider_id="unsupported", provider_type=ProviderType.ollama, name="unsupported"), MockProvider(),
|
||||
)
|
||||
with pytest.raises(ApiError) as caught:
|
||||
bind(rig, capability, provider_id=provider_id)
|
||||
assert (caught.value.status_code, caught.value.code) == (422, code)
|
||||
assert rig.service.configuration() == ModelRoutingConfig()
|
||||
assert rig.requests == []
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def api(monkeypatch, no_real_http, _isolate_data_dir):
|
||||
# Import the production container only after temporary storage is configured.
|
||||
from app import container as container_module, routes
|
||||
from app.main import app
|
||||
|
||||
containers = []
|
||||
|
||||
def restart():
|
||||
container = container_module.build_container()
|
||||
container.model_routing.credentials = FakeCredentials()
|
||||
|
||||
def unexpected(request):
|
||||
pytest.fail(f"Unexpected API-side provider HTTP: {request.url}")
|
||||
|
||||
container.model_routing.transport = httpx.MockTransport(unexpected)
|
||||
monkeypatch.setattr(container_module, "container", container)
|
||||
monkeypatch.setattr(routes, "container", container)
|
||||
containers.append(container)
|
||||
return container
|
||||
|
||||
container = restart()
|
||||
client = TestClient(app)
|
||||
yield SimpleNamespace(client=client, container=container, restart=restart)
|
||||
client.close()
|
||||
for container in containers:
|
||||
container.plugins.shutdown()
|
||||
container.mcp_servers.shutdown()
|
||||
|
||||
|
||||
def create_api_provider(api):
|
||||
result = api.client.post("/api/providers", json={
|
||||
"provider_type": "openai_compatible", "name": "Persisted fake",
|
||||
"base_url": "https://persist.invalid/v1", "default_model": "fake-model",
|
||||
})
|
||||
assert result.status_code == 200, result.text
|
||||
return result.json()
|
||||
|
||||
|
||||
def test_api_config_conflict_reference_delete_and_restart_persistence(api):
|
||||
provider = create_api_provider(api)
|
||||
provider_id = provider["provider_id"]
|
||||
assert api.client.get("/api/model-routing").json()["config"]["version"] == 0
|
||||
config = {"version": 0, "embedding": {"provider_id": provider_id, "model": "embed-model", "endpoint": "/embeddings"}}
|
||||
saved = api.client.put("/api/model-routing", json=config)
|
||||
assert saved.status_code == 200
|
||||
assert saved.json()["config"]["version"] == 1
|
||||
conflict = api.client.put("/api/model-routing", json=config)
|
||||
assert conflict.status_code == 409
|
||||
assert conflict.json()["error"]["code"] == "MODEL_ROUTING_VERSION_CONFLICT"
|
||||
blocked = api.client.delete(f"/api/providers/{provider_id}")
|
||||
assert blocked.status_code == 409 and blocked.json()["error"]["code"] == "PROVIDER_IN_USE"
|
||||
restarted = api.restart()
|
||||
assert restarted.providers.get_any(provider_id).config.model_dump(mode="json") == provider
|
||||
assert api.client.get("/api/model-routing").json()["config"] == saved.json()["config"]
|
||||
assert {item["provider_id"] for item in api.client.get("/api/providers").json()["items"]} == {"mock", provider_id}
|
||||
cleared = api.client.put("/api/model-routing", json={"version": 1})
|
||||
assert cleared.status_code == 200
|
||||
assert api.client.delete(f"/api/providers/{provider_id}").status_code == 200
|
||||
api.restart()
|
||||
assert api.client.get(f"/api/providers/{provider_id}").status_code == 404
|
||||
assert api.client.get("/api/model-routing").json()["config"]["version"] == 2
|
||||
|
||||
|
||||
def test_api_provider_type_patch_rebuilds_adapter_and_persists(api):
|
||||
from app.providers.anthropic_messages import AnthropicMessagesProvider
|
||||
|
||||
provider = create_api_provider(api)
|
||||
provider_id = provider["provider_id"]
|
||||
changed = api.client.patch(f"/api/providers/{provider_id}", json={
|
||||
"provider_type": "anthropic_messages", "base_url": "https://anthropic.invalid/v1",
|
||||
})
|
||||
assert changed.status_code == 200, changed.text
|
||||
assert changed.json()["provider_type"] == "anthropic_messages"
|
||||
assert changed.json()["name"] == provider["name"]
|
||||
assert isinstance(api.container.providers.get_any(provider_id).adapter, AnthropicMessagesProvider)
|
||||
restarted = api.restart()
|
||||
assert isinstance(restarted.providers.get_any(provider_id).adapter, AnthropicMessagesProvider)
|
||||
assert api.client.get(f"/api/providers/{provider_id}").json() == changed.json()
|
||||
for invalid_type in (None, "mock", "nonexistent-type"):
|
||||
rejected = api.client.patch(f"/api/providers/{provider_id}", json={"provider_type": invalid_type})
|
||||
assert rejected.status_code == 422
|
||||
assert api.client.get(f"/api/providers/{provider_id}").json() == changed.json()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("endpoint", ["https://elsewhere.invalid/embed", "//elsewhere.invalid/embed", "relative", "/../embed", "/embed?key=test"])
|
||||
def test_api_config_rejects_non_provider_endpoint_paths(api, endpoint):
|
||||
provider = create_api_provider(api)
|
||||
result = api.client.put("/api/model-routing", json={
|
||||
"embedding": {"provider_id": provider["provider_id"], "model": "embed", "endpoint": endpoint},
|
||||
})
|
||||
assert result.status_code == 422
|
||||
assert api.client.get("/api/model-routing").json()["config"]["version"] == 0
|
||||
|
||||
|
||||
def test_api_embedding_reports_remote_and_fallback_sources(api):
|
||||
provider = create_api_provider(api)
|
||||
assert api.client.put("/api/model-routing", json={
|
||||
"embedding": {"provider_id": provider["provider_id"], "model": "embed", "endpoint": "/embeddings"},
|
||||
}).status_code == 200
|
||||
api.container.model_routing.transport = httpx.MockTransport(
|
||||
lambda request: response({"data": [{"index": 0, "embedding": [3, 4]}]}),
|
||||
)
|
||||
result = api.client.post("/api/models/embeddings", json={"texts": ["hello"]})
|
||||
assert result.status_code == 200
|
||||
assert result.json()["source"] == "api" and result.json()["vectors"][0] == pytest.approx([0.6, 0.8])
|
||||
api.container.model_routing.transport = httpx.MockTransport(lambda request: response({"error": "denied"}, 401))
|
||||
result = api.client.post("/api/models/embeddings", json={"texts": ["hello"]})
|
||||
assert result.status_code == 200
|
||||
assert result.json()["source"] == "local" and result.json()["model_id"] == "hash-v1"
|
||||
assert result.json()["fallback_reason"] == "PROVIDER_AUTH_FAILED"
|
||||
assert api.client.post("/api/models/embeddings", json={"texts": []}).status_code == 422
|
||||
|
||||
|
||||
def test_api_speech_failure_reports_reason_in_503_and_transcription_job(api):
|
||||
from app.services.attachment_service import attachment_path
|
||||
|
||||
source, reference = attachment_path("audio.wav"), attachment_path("reference.wav")
|
||||
source.parent.mkdir(parents=True, exist_ok=True)
|
||||
source.write_bytes(b"test audio")
|
||||
reference.write_bytes(b"test reference")
|
||||
provider = create_api_provider(api)
|
||||
assert api.client.put("/api/model-routing", json={
|
||||
"transcription": {"provider_id": provider["provider_id"], "model": "asr", "endpoint": "/audio/transcriptions"},
|
||||
"speaker_matching": {"provider_id": provider["provider_id"], "model": "voice", "endpoint": "/audio/speaker-matches"},
|
||||
}).status_code == 200
|
||||
api.container.model_routing.transport = httpx.MockTransport(lambda request: response({"error": "offline"}, 500))
|
||||
match = api.client.post("/api/media/speaker-matches", json={"attachment_id": source.name, "reference_attachment_id": reference.name})
|
||||
assert match.status_code == 503
|
||||
assert match.json()["error"]["code"] == "LOCAL_MODEL_NOT_INSTALLED"
|
||||
assert match.json()["error"]["details"] == {"fallback_reason": "PROVIDER_UNAVAILABLE"}
|
||||
transcript = api.client.post("/api/media/transcriptions", json={"attachment_id": source.name, "language": "zh"})
|
||||
assert transcript.status_code == 202
|
||||
job = transcript.json()
|
||||
assert job["status"] == "failed" and job["error_code"] == "LOCAL_MODEL_NOT_INSTALLED"
|
||||
assert job["fallback_reason"] == "PROVIDER_UNAVAILABLE"
|
||||
assert api.client.get(f"/api/media/transcriptions/{job['job_id']}").json() == job
|
||||
|
||||
|
||||
@pytest.mark.parametrize("capability", ["embedding", "speaker_matching"])
|
||||
def test_out_of_float_range_json_number_is_invalid_remote_and_falls_back(rig, audio, capability):
|
||||
"""JSON integers may be finite but too large to convert to a Python float."""
|
||||
bind(rig, capability)
|
||||
data = {"data": [{"index": 0, "embedding": [10 ** 400, 1]}]} if capability == "embedding" else {"score": 10 ** 400}
|
||||
rig.http.handler = lambda request: response(data)
|
||||
if capability == "embedding":
|
||||
assert_local(rig, run(rig.service.embed(["text"])), ["text"], "PROVIDER_INVALID_RESPONSE")
|
||||
else:
|
||||
result = run(media_call(rig, capability, audio))
|
||||
assert result.source == "local" and result.score == rig.speech.score
|
||||
assert result.fallback_reason == "PROVIDER_INVALID_RESPONSE"
|
||||
@@ -84,7 +84,7 @@ def test_openai_compatible_maps_tool_call_and_credentials() -> None:
|
||||
)
|
||||
)
|
||||
|
||||
assert captured["tools"][0]["function"]["name"] == "math.add"
|
||||
assert captured["tools"][0]["function"]["name"].startswith("tool_")
|
||||
assert turn.tool_calls[0].name == "math.add"
|
||||
assert turn.tool_calls[0].arguments == {"left": 1, "right": 2}
|
||||
assert turn.input_tokens == 8
|
||||
|
||||
@@ -0,0 +1,610 @@
|
||||
"""Wire-level provider tests: no credentials, SDKs, clocks, or network services."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from app.contracts import Message, MessageRole, ModelCapability, ModelEventType as E, ModelRequest, ToolCall, ToolDefinition
|
||||
from app.providers.anthropic_messages import AnthropicMessagesProvider
|
||||
from app.providers.base import ProviderError
|
||||
from app.providers.ollama import OllamaProvider
|
||||
from app.providers.openai_compatible import OpenAICompatibleProvider
|
||||
from app.providers.openai_responses import OpenAIResponsesProvider
|
||||
|
||||
|
||||
NATIVE = ["responses", "anthropic"]
|
||||
PROTOCOLS = [*NATIVE, "compatible", "ollama"]
|
||||
SECRET = "test-only-sensitive-upstream-body"
|
||||
|
||||
|
||||
class Credentials:
|
||||
def resolve(self, credential_id):
|
||||
return SECRET if credential_id else None
|
||||
|
||||
|
||||
class Bytes(httpx.AsyncByteStream):
|
||||
def __init__(self, body: bytes, *, fragment: int = 17):
|
||||
self.body = body
|
||||
self.fragment = fragment
|
||||
self.closed = False
|
||||
|
||||
async def __aiter__(self):
|
||||
for offset in range(0, len(self.body), self.fragment):
|
||||
yield self.body[offset:offset + self.fragment]
|
||||
|
||||
async def aclose(self):
|
||||
self.closed = True
|
||||
|
||||
|
||||
class GatedBytes(Bytes):
|
||||
def __init__(self, body):
|
||||
super().__init__(body)
|
||||
self.waiting = asyncio.Event()
|
||||
self.release = asyncio.Event()
|
||||
|
||||
async def __aiter__(self):
|
||||
yield self.body
|
||||
self.waiting.set()
|
||||
await self.release.wait()
|
||||
|
||||
|
||||
def provider(protocol, handler, *, credential_id="test"):
|
||||
transport = httpx.MockTransport(handler)
|
||||
if protocol == "ollama":
|
||||
return OllamaProvider("https://provider.test", transport=transport)
|
||||
cls = {"responses": OpenAIResponsesProvider, "anthropic": AnthropicMessagesProvider,
|
||||
"compatible": OpenAICompatibleProvider}[protocol]
|
||||
return cls("https://provider.test/v1/", credential_id, Credentials(), transport=transport)
|
||||
|
||||
|
||||
def request(*, history=False):
|
||||
messages = [Message(role=MessageRole.user, content="查笔记")]
|
||||
if history:
|
||||
messages += [
|
||||
Message(role=MessageRole.system, content="Additional rules"),
|
||||
Message(role=MessageRole.assistant, content="Checking", tool_calls=[
|
||||
ToolCall(tool_call_id="old_1", name="lookup", arguments={"query": "a"}),
|
||||
ToolCall(tool_call_id="old_2", name="lookup", arguments={"query": "b"}),
|
||||
]),
|
||||
Message(role=MessageRole.tool, tool_call_id="old_1", content='{"found":1}'),
|
||||
Message(role=MessageRole.tool, tool_call_id="old_2", content='{"found":2}'),
|
||||
]
|
||||
return ModelRequest(
|
||||
provider_id="test", model="model", system="System rules", messages=messages,
|
||||
tools=[ToolDefinition(name="lookup", description="Find notes", parameters={"type": "object"})],
|
||||
max_tokens=512, temperature=0,
|
||||
)
|
||||
|
||||
|
||||
async def collect(iterator):
|
||||
return [event async for event in iterator]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name", ["lookup", "notes.search"])
|
||||
def test_compatible_split_tool_name_preserves_identity(name):
|
||||
from app.providers.tool_names import prepare_tool_names
|
||||
req = request()
|
||||
req.tools[0].name = name
|
||||
wire, _ = prepare_tool_names(req)
|
||||
alias = wire.tools[0].name
|
||||
|
||||
def handler(_):
|
||||
return httpx.Response(200, content=sse(
|
||||
{"choices": [{"delta": {"tool_calls": [{"index": 0, "id": "call_1",
|
||||
"function": {"name": alias[:3], "arguments": ""}}]}}]},
|
||||
{"choices": [{"delta": {"tool_calls": [{"index": 0,
|
||||
"function": {"name": alias[3:], "arguments": '{"query":"x"}'}}]},
|
||||
"finish_reason": "tool_calls"}]},
|
||||
{"type": "[DONE]"},
|
||||
))
|
||||
|
||||
events = asyncio.run(collect(provider("compatible", handler).stream(req)))
|
||||
assert [e.data["name"] for e in events if e.event == E.tool_call_start] == [name]
|
||||
assert json.loads("".join(e.data["arguments_delta"] for e in events
|
||||
if e.event == E.tool_call_delta)) == {"query": "x"}
|
||||
assert events[-1].data["status"] == "completed"
|
||||
|
||||
|
||||
def sse(*events):
|
||||
return "".join(
|
||||
f"event: {event.get('type', 'message')}\r\ndata: {json.dumps(event, ensure_ascii=False)}\r\n\r\n"
|
||||
for event in events
|
||||
).encode()
|
||||
|
||||
|
||||
def wire(protocol, *events):
|
||||
if protocol == "ollama":
|
||||
return ("\n".join(json.dumps(event, ensure_ascii=False) for event in events) + "\n").encode()
|
||||
return sse(*events)
|
||||
|
||||
|
||||
def start(protocol):
|
||||
if protocol == "responses":
|
||||
return [{"type": "response.output_text.delta", "delta": "你好"}]
|
||||
if protocol == "anthropic":
|
||||
return [{"type": "message_start", "message": {"usage": {"input_tokens": 7, "output_tokens": 0}}},
|
||||
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "你好"}}]
|
||||
if protocol == "compatible":
|
||||
return [{"choices": [{"delta": {"content": "你好"}}]}]
|
||||
return [{"message": {"content": "你好"}, "done": False}]
|
||||
|
||||
|
||||
def terminal(protocol):
|
||||
if protocol == "responses":
|
||||
return [{"type": "response.completed", "response": {"status": "completed", "usage": {"input_tokens": 7, "output_tokens": 2}}}]
|
||||
if protocol == "anthropic":
|
||||
return [{"type": "content_block_stop", "index": 0},
|
||||
{"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 2}},
|
||||
{"type": "message_stop"}]
|
||||
if protocol == "compatible":
|
||||
return [{"choices": [{"delta": {}, "finish_reason": "stop"}], "usage": {"prompt_tokens": 7, "completion_tokens": 2}}]
|
||||
return [{"message": {}, "done": True, "prompt_eval_count": 7, "eval_count": 2}]
|
||||
|
||||
|
||||
def assert_events(events):
|
||||
assert events[-1].event == E.done
|
||||
assert events[-1].data["status"] == ("failed" if any(event.event == E.error for event in events) else "completed")
|
||||
assert sum(event.event == E.done for event in events) == 1
|
||||
assert [event.sequence for event in events] == list(range(len(events)))
|
||||
assert all(event.timestamp.tzinfo is not None for event in events)
|
||||
|
||||
|
||||
def assert_error(events, code):
|
||||
assert_events(events)
|
||||
assert events[-2].event == E.error
|
||||
assert events[-2].data["code"] == code
|
||||
assert SECRET not in str(events[-2].data)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("protocol", NATIVE)
|
||||
def test_native_completion_and_history(protocol):
|
||||
captured = {}
|
||||
|
||||
def handler(req):
|
||||
captured.update(json.loads(req.content))
|
||||
assert req.url.path == ("/v1/responses" if protocol == "responses" else "/v1/messages")
|
||||
if protocol == "responses":
|
||||
assert req.headers["authorization"] == f"Bearer {SECRET}"
|
||||
body = {"status": "completed", "output": [
|
||||
{"type": "reasoning", "summary": [{"type": "summary_text", "text": "thinking"}]},
|
||||
{"type": "message", "content": [{"type": "output_text", "text": "完成"}]},
|
||||
{"type": "function_call", "call_id": "next", "name": "lookup", "arguments": '{"query":"c"}'},
|
||||
], "usage": {"input_tokens": 10, "output_tokens": 3}}
|
||||
else:
|
||||
assert "authorization" not in req.headers
|
||||
assert req.headers["x-api-key"] == SECRET
|
||||
assert req.headers["anthropic-version"] == "2023-06-01"
|
||||
body = {"type": "message", "content": [
|
||||
{"type": "thinking", "thinking": "thinking", "signature": "sig"},
|
||||
{"type": "text", "text": "完成"},
|
||||
{"type": "tool_use", "id": "next", "name": "lookup", "input": {"query": "c"}},
|
||||
], "usage": {"input_tokens": 5, "cache_creation_input_tokens": 2, "cache_read_input_tokens": 3, "output_tokens": 3}}
|
||||
return httpx.Response(200, json=body)
|
||||
|
||||
turn = asyncio.run(provider(protocol, handler).complete(request(history=True)))
|
||||
assert turn.text == "完成"
|
||||
assert (turn.input_tokens, turn.output_tokens) == (10, 3)
|
||||
assert turn.tool_calls[0].tool_call_id == "next"
|
||||
assert turn.tool_calls[0].arguments == {"query": "c"}
|
||||
assert captured["stream"] is False
|
||||
assert captured["temperature"] == 0
|
||||
if protocol == "responses":
|
||||
assert captured["instructions"] == "System rules"
|
||||
assert captured["max_output_tokens"] == 512
|
||||
assert captured["tools"][0]["parameters"] == {"type": "object"}
|
||||
calls = [item for item in captured["input"] if item.get("type") == "function_call"]
|
||||
outputs = [item for item in captured["input"] if item.get("type") == "function_call_output"]
|
||||
assert [call["call_id"] for call in calls] == ["old_1", "old_2"]
|
||||
assert json.loads(calls[1]["arguments"]) == {"query": "b"}
|
||||
assert outputs == [{"type": "function_call_output", "call_id": "old_1", "output": '{"found":1}'},
|
||||
{"type": "function_call_output", "call_id": "old_2", "output": '{"found":2}'}]
|
||||
assert {"role": "system", "content": "Additional rules"} in captured["input"]
|
||||
else:
|
||||
assert captured["system"] == "System rules\n\nAdditional rules"
|
||||
assert captured["max_tokens"] == 512
|
||||
assert captured["tools"][0]["input_schema"] == {"type": "object"}
|
||||
assert captured["messages"][1]["content"][2] == {
|
||||
"type": "tool_use", "id": "old_2", "name": "lookup", "input": {"query": "b"},
|
||||
}
|
||||
assert captured["messages"][-1] == {"role": "user", "content": [
|
||||
{"type": "tool_result", "tool_use_id": "old_1", "content": '{"found":1}'},
|
||||
{"type": "tool_result", "tool_use_id": "old_2", "content": '{"found":2}'},
|
||||
]}
|
||||
|
||||
|
||||
def responses_tool_events():
|
||||
events = [
|
||||
{"type": "response.created", "response": {"usage": {"input_tokens": 10, "output_tokens": 0}}},
|
||||
{"type": "response.reasoning_summary_text.delta", "delta": "计划"},
|
||||
{"type": "response.output_text.delta", "delta": "查"},
|
||||
{"type": "response.output_text.delta", "delta": "找"},
|
||||
]
|
||||
for index in (2, 3):
|
||||
events.append({"type": "response.output_item.added", "output_index": index, "item": {
|
||||
"id": f"item_{index}", "type": "function_call", "call_id": f"call_{index}", "name": "lookup", "arguments": "",
|
||||
}})
|
||||
for index, fragment in [(2, '{"query":'), (3, '{}'), (2, '"笔记"}')]:
|
||||
events.append({"type": "response.function_call_arguments.delta", "output_index": index,
|
||||
"item_id": f"item_{index}", "delta": fragment})
|
||||
for index, arguments in [(3, '{}'), (2, '{"query":"笔记"}')]:
|
||||
events += [
|
||||
{"type": "response.function_call_arguments.done", "output_index": index, "item_id": f"item_{index}", "arguments": arguments},
|
||||
{"type": "response.output_item.done", "output_index": index, "item": {
|
||||
"id": f"item_{index}", "type": "function_call", "call_id": f"call_{index}", "name": "lookup", "arguments": arguments,
|
||||
}},
|
||||
]
|
||||
events += [{"type": "future.event"}, {"type": "response.completed", "response": {
|
||||
"status": "completed", "usage": {"input_tokens": 10, "output_tokens": 9},
|
||||
}}]
|
||||
return events
|
||||
|
||||
|
||||
def anthropic_tool_events():
|
||||
events = [
|
||||
{"type": "message_start", "message": {"usage": {
|
||||
"input_tokens": 5, "cache_read_input_tokens": 3, "cache_creation_input_tokens": 2, "output_tokens": 1,
|
||||
}}},
|
||||
{"type": "ping"},
|
||||
{"type": "content_block_start", "index": 0, "content_block": {"type": "thinking", "thinking": ""}},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "thinking_delta", "thinking": "计划"}},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "signature_delta", "signature": "sig"}},
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
{"type": "content_block_start", "index": 1, "content_block": {"type": "text", "text": "查"}},
|
||||
{"type": "content_block_delta", "index": 1, "delta": {"type": "text_delta", "text": "找"}},
|
||||
{"type": "content_block_stop", "index": 1},
|
||||
]
|
||||
for index, fragments in [(2, ['{"query":', '"笔记"}']), (3, [])]:
|
||||
events.append({"type": "content_block_start", "index": index, "content_block": {
|
||||
"type": "tool_use", "id": f"call_{index}", "name": "lookup", "input": {},
|
||||
}})
|
||||
for fragment in fragments:
|
||||
events.append({"type": "content_block_delta", "index": index,
|
||||
"delta": {"type": "input_json_delta", "partial_json": fragment}})
|
||||
events.append({"type": "content_block_stop", "index": index})
|
||||
events += [
|
||||
{"type": "message_delta", "delta": {"stop_reason": "tool_use"}, "usage": {"output_tokens": 4}},
|
||||
{"type": "future.event"},
|
||||
{"type": "message_delta", "delta": {}, "usage": {"output_tokens": 9}},
|
||||
{"type": "message_stop"},
|
||||
]
|
||||
return events
|
||||
|
||||
|
||||
@pytest.mark.parametrize("protocol", NATIVE)
|
||||
def test_native_stream_tools_reasoning_usage_and_fragmented_utf8(protocol):
|
||||
frames = responses_tool_events() if protocol == "responses" else anthropic_tool_events()
|
||||
body = Bytes(b": comment\r\n\r\n" + sse(*frames) + b"data: malformed after completion\n\n", fragment=1)
|
||||
|
||||
def handler(req):
|
||||
payload = json.loads(req.content)
|
||||
assert payload["stream"] is True
|
||||
assert payload["tools"]
|
||||
assert (payload.get("input") or payload.get("messages"))
|
||||
return httpx.Response(200, stream=body)
|
||||
|
||||
events = asyncio.run(collect(provider(protocol, handler).stream(request(history=True))))
|
||||
assert_events(events)
|
||||
assert not any(event.event == E.error for event in events)
|
||||
assert [event.data["text"] for event in events if event.event == E.text_delta] == ["查", "找"]
|
||||
assert [event.data["text"] for event in events if event.event == E.thinking_delta] == ["计划"]
|
||||
assert [event.data["tool_call_id"] for event in events if event.event == E.tool_call_start] == ["call_2", "call_3"]
|
||||
assert sorted(event.data["tool_call_id"] for event in events if event.event == E.tool_call_end) == ["call_2", "call_3"]
|
||||
for call_id, expected in [("call_2", {"query": "笔记"}), ("call_3", {})]:
|
||||
arguments = "".join(event.data["arguments_delta"] for event in events
|
||||
if event.event == E.tool_call_delta and event.data["tool_call_id"] == call_id)
|
||||
assert json.loads(arguments) == expected
|
||||
usages = [event.data for event in events if event.event == E.usage]
|
||||
assert usages[-1] == {"input_tokens": 10, "output_tokens": 9, "total_tokens": 19}
|
||||
assert all(usage["input_tokens"] == 10 for usage in usages)
|
||||
if protocol == "anthropic":
|
||||
assert [usage["output_tokens"] for usage in usages] == [1, 4, 9]
|
||||
assert body.closed
|
||||
|
||||
|
||||
@pytest.mark.parametrize("protocol", PROTOCOLS)
|
||||
def test_stream_terminal_usage_and_closure(protocol):
|
||||
body = Bytes(wire(protocol, *start(protocol), *terminal(protocol)))
|
||||
events = asyncio.run(collect(provider(protocol, lambda _: httpx.Response(200, stream=body)).stream(request())))
|
||||
assert_events(events)
|
||||
assert not any(event.event == E.error for event in events)
|
||||
assert [event.data["text"] for event in events if event.event == E.text_delta] == ["你好"]
|
||||
assert [event.data for event in events if event.event == E.usage][-1] == {"input_tokens": 7, "output_tokens": 2, "total_tokens": 9}
|
||||
assert body.closed
|
||||
|
||||
|
||||
@pytest.mark.parametrize("protocol", PROTOCOLS)
|
||||
@pytest.mark.parametrize("empty", [False, True])
|
||||
def test_truncated_stream(protocol, empty):
|
||||
body = Bytes(b"" if empty else wire(protocol, *start(protocol)))
|
||||
events = asyncio.run(collect(provider(protocol, lambda _: httpx.Response(200, stream=body)).stream(request())))
|
||||
assert_error(events, "PROVIDER_STREAM_TRUNCATED")
|
||||
assert body.closed
|
||||
|
||||
|
||||
@pytest.mark.parametrize("protocol", PROTOCOLS)
|
||||
@pytest.mark.parametrize("bad", [b"not-json", b"[]", b"null", b'{"usage":'])
|
||||
def test_malformed_stream_is_sanitized(protocol, bad):
|
||||
suffix = bad + b"\n" if protocol == "ollama" else b"data: " + bad + b"\n\n"
|
||||
body = Bytes(wire(protocol, *start(protocol)) + suffix)
|
||||
events = asyncio.run(collect(provider(protocol, lambda _: httpx.Response(200, stream=body)).stream(request())))
|
||||
assert_error(events, "PROVIDER_INVALID_RESPONSE")
|
||||
assert body.closed
|
||||
|
||||
|
||||
@pytest.mark.parametrize("protocol", PROTOCOLS)
|
||||
@pytest.mark.parametrize("error_type,code", [("rate_limit_error", "PROVIDER_RATE_LIMITED"),
|
||||
("authentication_error", "PROVIDER_AUTH_FAILED"),
|
||||
("overloaded_error", "PROVIDER_UNAVAILABLE")])
|
||||
def test_in_band_error_after_partial_output(protocol, error_type, code):
|
||||
body = Bytes(wire(protocol, *start(protocol), {"type": "error", "error": {"type": error_type, "message": SECRET}}))
|
||||
events = asyncio.run(collect(provider(protocol, lambda _: httpx.Response(200, stream=body)).stream(request())))
|
||||
assert any(event.event == E.text_delta for event in events)
|
||||
assert_error(events, code)
|
||||
assert body.closed
|
||||
|
||||
|
||||
@pytest.mark.parametrize("protocol", PROTOCOLS)
|
||||
@pytest.mark.parametrize("status,code", [(400, "PROVIDER_INVALID_REQUEST"), (401, "PROVIDER_AUTH_FAILED"),
|
||||
(403, "PROVIDER_AUTH_FAILED"), (404, "MODEL_NOT_FOUND"),
|
||||
(429, "PROVIDER_RATE_LIMITED"), (500, "PROVIDER_UNAVAILABLE")])
|
||||
def test_http_errors_completion_and_stream(protocol, status, code):
|
||||
adapter = provider(protocol, lambda _: httpx.Response(status, text=SECRET))
|
||||
with pytest.raises(ProviderError) as exc:
|
||||
asyncio.run(adapter.complete(request()))
|
||||
assert exc.value.code == code
|
||||
assert SECRET not in str(exc.value)
|
||||
assert_error(asyncio.run(collect(adapter.stream(request()))), code)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("protocol", PROTOCOLS)
|
||||
@pytest.mark.parametrize("body,code", [(b"broken", "PROVIDER_INVALID_RESPONSE"),
|
||||
(b"[]", "PROVIDER_INVALID_RESPONSE"),
|
||||
(b"{}", "PROVIDER_INVALID_RESPONSE"),
|
||||
(json.dumps({"error": {"code": "invalid_api_key", "message": SECRET}}).encode(), "PROVIDER_AUTH_FAILED")])
|
||||
def test_bad_completion(protocol, body, code):
|
||||
adapter = provider(protocol, lambda _: httpx.Response(200, content=body))
|
||||
with pytest.raises(ProviderError) as exc:
|
||||
asyncio.run(adapter.complete(request()))
|
||||
assert exc.value.code == code
|
||||
assert SECRET not in str(exc.value)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("protocol", PROTOCOLS)
|
||||
@pytest.mark.parametrize("error,code", [(httpx.ReadTimeout, "PROVIDER_TIMEOUT"),
|
||||
(httpx.ConnectError, "PROVIDER_UNAVAILABLE")])
|
||||
def test_transport_error_mapping(protocol, error, code):
|
||||
def handler(req):
|
||||
raise error(SECRET, request=req)
|
||||
|
||||
adapter = provider(protocol, handler)
|
||||
with pytest.raises(ProviderError) as exc:
|
||||
asyncio.run(adapter.complete(request()))
|
||||
assert exc.value.code == code
|
||||
assert SECRET not in str(exc.value)
|
||||
assert_error(asyncio.run(collect(adapter.stream(request()))), code)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("protocol", PROTOCOLS)
|
||||
@pytest.mark.parametrize("cancel", [True, False])
|
||||
def test_incremental_delivery_cancellation_and_explicit_close(protocol, cancel):
|
||||
async def scenario():
|
||||
body = GatedBytes(wire(protocol, *start(protocol)))
|
||||
adapter = provider(protocol, lambda _: httpx.Response(200, stream=body))
|
||||
iterator = adapter.stream(request())
|
||||
seen = []
|
||||
while True:
|
||||
event = await asyncio.wait_for(anext(iterator), timeout=1)
|
||||
seen.append(event)
|
||||
if event.event == E.text_delta:
|
||||
break
|
||||
# The first token arrives while the response is still open and blocked.
|
||||
assert seen[-1].data["text"] == "你好"
|
||||
assert not body.closed
|
||||
if cancel:
|
||||
pending = asyncio.create_task(anext(iterator))
|
||||
await asyncio.wait_for(body.waiting.wait(), timeout=1)
|
||||
pending.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await pending
|
||||
else:
|
||||
await iterator.aclose()
|
||||
assert body.closed
|
||||
assert not any(event.event in {E.error, E.done} for event in seen)
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("protocol", NATIVE)
|
||||
def test_cancellation_before_response_headers(protocol):
|
||||
async def scenario():
|
||||
entered = asyncio.Event()
|
||||
closed = asyncio.Event()
|
||||
|
||||
async def handler(req):
|
||||
entered.set()
|
||||
try:
|
||||
await asyncio.Event().wait()
|
||||
finally:
|
||||
closed.set()
|
||||
|
||||
adapter = provider(protocol, handler)
|
||||
pending = asyncio.create_task(adapter.complete(request()))
|
||||
await asyncio.wait_for(entered.wait(), timeout=1)
|
||||
pending.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await pending
|
||||
assert closed.is_set()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("protocol", NATIVE)
|
||||
def test_native_discovery_does_not_claim_non_chat_capabilities(protocol):
|
||||
def handler(req):
|
||||
assert req.url.path == "/v1/models"
|
||||
return httpx.Response(200, json={"data": [{"id": name} for name in ["chat-model", "text-embedding-3-small", "whisper-1", "gpt-audio"]]})
|
||||
|
||||
models = asyncio.run(provider(protocol, handler).list_models())
|
||||
assert ModelCapability.chat in models[0].capabilities
|
||||
assert models[1].capabilities == [ModelCapability.embedding]
|
||||
assert all(ModelCapability.chat not in model.capabilities for model in models[1:])
|
||||
|
||||
|
||||
@pytest.mark.parametrize("protocol", NATIVE)
|
||||
def test_native_structured_format_mapping(protocol):
|
||||
adapter = provider(protocol, lambda _: pytest.fail("No network expected"))
|
||||
req = request()
|
||||
req.response_format = {"type": "json_schema", "json_schema": {
|
||||
"name": "answer", "strict": True, "schema": {"type": "object", "properties": {}},
|
||||
}}
|
||||
payload = adapter._payload(req, stream=False)
|
||||
format_ = payload["text"]["format"] if protocol == "responses" else payload["output_config"]["format"]
|
||||
assert format_["type"] == "json_schema"
|
||||
assert format_["schema"] == {"type": "object", "properties": {}}
|
||||
if protocol == "responses":
|
||||
assert format_["name"] == "answer"
|
||||
assert format_["strict"] is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize("protocol", NATIVE)
|
||||
def test_invalid_tool_arguments_and_unclosed_tool(protocol):
|
||||
frames = responses_tool_events() if protocol == "responses" else anthropic_tool_events()
|
||||
# A syntactically valid terminal cannot rescue an unfinished tool block.
|
||||
index = next(i for i, frame in enumerate(frames)
|
||||
if frame["type"] in {"response.function_call_arguments.delta", "content_block_delta"}
|
||||
and (frame.get("output_index") == 2 or frame.get("index") == 2))
|
||||
partial = frames[:index + 1]
|
||||
final = frames[-1]
|
||||
events = asyncio.run(collect(provider(protocol, lambda _: httpx.Response(200, content=sse(*partial, final))).stream(request())))
|
||||
assert_error(events, "PROVIDER_STREAM_TRUNCATED")
|
||||
assert not any(event.event == E.tool_call_end for event in events)
|
||||
|
||||
for frame in frames:
|
||||
if frame["type"] == "response.function_call_arguments.done":
|
||||
frame["arguments"] = "[]"
|
||||
break
|
||||
if frame["type"] == "content_block_delta" and frame.get("index") == 2:
|
||||
frame["delta"]["partial_json"] = "malformed"
|
||||
break
|
||||
events = asyncio.run(collect(provider(protocol, lambda _: httpx.Response(200, content=sse(*frames))).stream(request())))
|
||||
assert_error(events, "PROVIDER_INVALID_RESPONSE")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("kind,code", [("response.failed", "PROVIDER_UNAVAILABLE"),
|
||||
("response.incomplete", "PROVIDER_INCOMPLETE_RESPONSE")])
|
||||
def test_responses_failed_and_incomplete(kind, code):
|
||||
frame = {"type": kind, "response": {"status": kind.split(".")[1], "incomplete_details": {"reason": SECRET}}}
|
||||
events = asyncio.run(collect(provider("responses", lambda _: httpx.Response(200, content=sse(*start("responses"), frame))).stream(request())))
|
||||
assert_error(events, code)
|
||||
|
||||
|
||||
def test_sse_multiline_data_and_event_name_without_json_type():
|
||||
body = (b': keepalive\n\nevent: response.output_text.delta\ndata: {\ndata: "delta": "hello"\ndata: }\n\n'
|
||||
+ sse({"type": "response.completed", "response": {"status": "completed"}}))
|
||||
events = asyncio.run(collect(provider("responses", lambda _: httpx.Response(200, content=body)).stream(request())))
|
||||
assert_events(events)
|
||||
assert [event.data["text"] for event in events if event.event == E.text_delta] == ["hello"]
|
||||
assert not any(event.event == E.error for event in events)
|
||||
|
||||
|
||||
def test_ollama_history_options_and_in_band_string_error():
|
||||
captured = {}
|
||||
|
||||
def handler(req):
|
||||
captured.update(json.loads(req.content))
|
||||
return httpx.Response(200, json={"error": SECRET})
|
||||
|
||||
with pytest.raises(ProviderError) as exc:
|
||||
asyncio.run(provider("ollama", handler).complete(request(history=True)))
|
||||
assert exc.value.code == "PROVIDER_UNAVAILABLE"
|
||||
assert SECRET not in str(exc.value)
|
||||
assert captured["messages"][-1]["tool_name"] == "lookup"
|
||||
assert captured["options"] == {"temperature": 0.0, "num_predict": 512}
|
||||
|
||||
@pytest.mark.parametrize("protocol", PROTOCOLS)
|
||||
@pytest.mark.parametrize("streaming", [False, True])
|
||||
def test_namespaced_tools_roundtrip_without_changing_internal_request(protocol, streaming):
|
||||
import re
|
||||
model_request = request(history=True)
|
||||
original_name = "mcp.my-server.search.notes"
|
||||
model_request.tools[0].name = original_name
|
||||
for message in model_request.messages:
|
||||
for call in message.tool_calls:
|
||||
call.name = original_name
|
||||
before = model_request.model_dump()
|
||||
|
||||
def handler(req):
|
||||
payload = json.loads(req.content)
|
||||
definition = payload["tools"][0]
|
||||
name = (definition.get("function") or definition)["name"]
|
||||
assert name != original_name and re.fullmatch(r"[a-zA-Z0-9_-]{1,64}", name)
|
||||
assert original_name not in req.content.decode()
|
||||
if protocol == "responses":
|
||||
item = {"type": "function_call", "id": "item1", "call_id": "call1", "name": name, "arguments": "{}"}
|
||||
body = {"status": "completed", "output": [item]}
|
||||
events = [
|
||||
{"type": "response.output_item.done", "output_index": 0, "item": item},
|
||||
{"type": "response.completed", "response": {"status": "completed"}},
|
||||
]
|
||||
elif protocol == "anthropic":
|
||||
item = {"type": "tool_use", "id": "call1", "name": name, "input": {}}
|
||||
body = {"content": [item]}
|
||||
events = [
|
||||
{"type": "message_start", "message": {}},
|
||||
{"type": "content_block_start", "index": 0, "content_block": item},
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
{"type": "message_stop"},
|
||||
]
|
||||
elif protocol == "compatible":
|
||||
item = {"id": "call1", "function": {"name": name, "arguments": "{}"}}
|
||||
body = {"choices": [{"message": {"tool_calls": [item]}}]}
|
||||
events = [{"choices": [{"delta": {"tool_calls": [{"index": 0, **item}]}, "finish_reason": "tool_calls"}]}]
|
||||
else:
|
||||
item = {"function": {"name": name, "arguments": {}}}
|
||||
body = {"message": {"tool_calls": [item]}, "done": True}
|
||||
events = [body]
|
||||
return httpx.Response(200, content=wire(protocol, *events)) if streaming else httpx.Response(200, json=body)
|
||||
|
||||
adapter = provider(protocol, handler)
|
||||
if streaming:
|
||||
events = asyncio.run(collect(adapter.stream(model_request)))
|
||||
assert_events(events)
|
||||
assert [event.data["name"] for event in events if event.event == E.tool_call_start] == [original_name]
|
||||
else:
|
||||
assert asyncio.run(adapter.complete(model_request)).tool_calls[0].name == original_name
|
||||
assert model_request.model_dump() == before
|
||||
|
||||
def test_chat_route_closes_upstream_and_sanitizes_unexpected_errors(monkeypatch):
|
||||
from types import SimpleNamespace
|
||||
from datetime import datetime, timezone
|
||||
from app import routes
|
||||
from app.contracts import ChatRequest, ModelEvent
|
||||
closed = []
|
||||
|
||||
class Adapter:
|
||||
async def stream(self, request):
|
||||
try:
|
||||
yield ModelEvent(event=E.text_delta, sequence=0, data={"text": "first"}, timestamp=datetime.now(timezone.utc))
|
||||
raise RuntimeError(SECRET)
|
||||
finally:
|
||||
closed.append(True)
|
||||
|
||||
monkeypatch.setattr(routes, "provider_or_404", lambda _: SimpleNamespace(adapter=Adapter()))
|
||||
|
||||
async def scenario():
|
||||
response = await routes.chat(ChatRequest(provider_id="test", model="test", messages=[]))
|
||||
iterator = response.body_iterator
|
||||
await anext(iterator)
|
||||
await iterator.aclose()
|
||||
assert len(closed) == 1
|
||||
response = await routes.chat(ChatRequest(provider_id="test", model="test", messages=[]))
|
||||
items = [json.loads(chunk.split("data: ")[1].strip()) async for chunk in response.body_iterator]
|
||||
assert [item["sequence"] for item in items] == [0, 1, 2]
|
||||
assert items[-1]["data"]["status"] == "failed"
|
||||
assert SECRET not in str(items)
|
||||
assert len(closed) == 2
|
||||
|
||||
asyncio.run(scenario())
|
||||
@@ -436,6 +436,83 @@ def test_fts_pagination_is_not_truncated_at_one_thousand(vault) -> None:
|
||||
assert len(response.items) == 10
|
||||
|
||||
|
||||
def test_fts_score_threshold_filters_before_total(vault) -> None:
|
||||
"""score_threshold 先于计数与分页生效:total 反映过滤后数量,与 items 一致。
|
||||
|
||||
高阈值过滤掉全部结果时 total==0 且 items 为空,杜绝「空页但 total>0」的
|
||||
不一致(审阅 P2-7)。
|
||||
"""
|
||||
from app.retrieval.engine import engine
|
||||
from app.services import note_service
|
||||
|
||||
# 10 个 block,含「目标」次数递增,bm25 分数各异,min-max 归一化后分数落在 [0,1]
|
||||
markdown = "\n\n".join(f"{'目标' * i} 分隔内容" for i in range(1, 11))
|
||||
asyncio.run(
|
||||
note_service.create_note(title="阈值过滤", markdown=markdown, folder="", tags=[])
|
||||
)
|
||||
|
||||
all_hits = asyncio.run(
|
||||
engine.search(
|
||||
SearchRequest(query="目标", mode=SearchMode.fts, limit=20, score_threshold=0.0)
|
||||
)
|
||||
)
|
||||
filtered = asyncio.run(
|
||||
engine.search(
|
||||
SearchRequest(query="目标", mode=SearchMode.fts, limit=20, score_threshold=0.5)
|
||||
)
|
||||
)
|
||||
none = asyncio.run(
|
||||
engine.search(
|
||||
SearchRequest(query="目标", mode=SearchMode.fts, limit=20, score_threshold=2.0)
|
||||
)
|
||||
)
|
||||
|
||||
assert all_hits.page.total >= 10
|
||||
assert 0 < filtered.page.total < all_hits.page.total # 阈值过滤掉部分而非全部
|
||||
assert filtered.page.total == len(filtered.items)
|
||||
assert none.page.total == 0
|
||||
assert none.items == []
|
||||
|
||||
|
||||
def test_fts_offset_beyond_end_reports_real_total(vault) -> None:
|
||||
"""offset 越过末页时 items 为空,但 total 仍为真实命中数而非归零。"""
|
||||
from app.retrieval.engine import engine
|
||||
from app.services import note_service
|
||||
|
||||
asyncio.run(
|
||||
note_service.create_note(title="越界分页", markdown="检索 检索 检索 检索", folder="", tags=[])
|
||||
)
|
||||
|
||||
resp = asyncio.run(
|
||||
engine.search(SearchRequest(query="检索", mode=SearchMode.fts, limit=10, offset=100))
|
||||
)
|
||||
assert resp.page.total >= 1
|
||||
assert resp.items == []
|
||||
|
||||
|
||||
def test_fts_not_truncated_at_five_thousand(vault) -> None:
|
||||
"""FTS 结果不再被 5000 条上限截断:>5000 命中时 total 为真实计数,末页仍可访问。"""
|
||||
from app.retrieval.engine import engine
|
||||
from app.services import note_service
|
||||
|
||||
markdown = "\n\n".join(f"共同词 q{i}" for i in range(5010))
|
||||
asyncio.run(
|
||||
note_service.create_note(title="五千条分页", markdown=markdown, folder="", tags=[])
|
||||
)
|
||||
|
||||
first = asyncio.run(
|
||||
engine.search(SearchRequest(query="共同词", mode=SearchMode.fts, limit=10, offset=0))
|
||||
)
|
||||
assert first.page.total == 5010
|
||||
assert len(first.items) == 10
|
||||
|
||||
last = asyncio.run(
|
||||
engine.search(SearchRequest(query="共同词", mode=SearchMode.fts, limit=10, offset=5005))
|
||||
)
|
||||
assert last.page.total == 5010
|
||||
assert len(last.items) == 5
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 审阅回归:PATCH tags 语义 / 向量-块一致性 / 过滤漏召回 / rebuild 语义与回滚
|
||||
# --------------------------------------------------------------------------- #
|
||||
@@ -563,9 +640,10 @@ def test_rebuild_failure_restores_old_index(vault, monkeypatch) -> None:
|
||||
assert repository.stats() == before # 旧索引已恢复,无半成品
|
||||
|
||||
|
||||
def test_first_rebuild_failure_removes_partial_database(vault, monkeypatch) -> None:
|
||||
def test_first_rebuild_failure_leaves_no_partial_index(vault, monkeypatch) -> None:
|
||||
"""首次启动没有旧库时,失败也不能留下已经写入的部分索引。"""
|
||||
from app.services import index_service
|
||||
from app import repository
|
||||
|
||||
_write_vault(
|
||||
vault,
|
||||
@@ -574,17 +652,17 @@ def test_first_rebuild_failure_removes_partial_database(vault, monkeypatch) -> N
|
||||
real_index = index_service.index_note
|
||||
calls = {"count": 0}
|
||||
|
||||
async def fail_on_second(parsed):
|
||||
async def fail_on_second(parsed, **kwargs):
|
||||
calls["count"] += 1
|
||||
if calls["count"] == 2:
|
||||
raise RuntimeError("injected first-rebuild failure")
|
||||
await real_index(parsed)
|
||||
await real_index(parsed, **kwargs)
|
||||
|
||||
monkeypatch.setattr(index_service, "index_note", fail_on_second)
|
||||
with pytest.raises(RuntimeError):
|
||||
asyncio.run(index_service.rebuild(IndexRebuildRequest(scope="all")))
|
||||
|
||||
assert not get_settings().db_path.exists()
|
||||
assert repository.stats() == {"notes": 0, "blocks": 0}
|
||||
|
||||
|
||||
def test_rebuild_preserves_task_note_links(vault) -> None:
|
||||
|
||||
@@ -0,0 +1,466 @@
|
||||
"""Phase E route integration: deterministic runtimes, isolated DBs, no network."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from app import repository
|
||||
from app.config import get_settings
|
||||
from app.contracts import IndexRebuildRequest, SearchMode, SearchRequest
|
||||
from app.database.db import connect, transaction
|
||||
from app.retrieval import routed_vectors
|
||||
from app.retrieval.embedding import HashEmbeddingProvider
|
||||
from app.retrieval.engine import RetrievalEngine, engine
|
||||
from app.retrieval.reranker import LexicalReranker
|
||||
from app.retrieval.vectorstore import SqliteVecStore, VectorHit
|
||||
from app.services import index_service, note_service
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeRuntime:
|
||||
model_id: str = "space-a"
|
||||
dimensions: int = 3 # Deliberately differs from sqlite-vec's fixed 128.
|
||||
source: str = "api"
|
||||
error: BaseException | None = None
|
||||
calls: list[list[str]] = field(default_factory=list)
|
||||
result_override: object | None = None
|
||||
|
||||
async def embed(self, texts):
|
||||
self.calls.append(list(texts))
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
if self.result_override is not None:
|
||||
return self.result_override
|
||||
vectors = []
|
||||
for text in texts:
|
||||
# The API associates "apple" with banana; hash retrieval picks apple.
|
||||
first = text == "apple orchard"
|
||||
if self.model_id == "space-b":
|
||||
first = not first
|
||||
vectors.append(([1.0, 0.0] if first else [0.0, 1.0]) + [0.0] * (self.dimensions - 2))
|
||||
return SimpleNamespace(
|
||||
vectors=vectors, source=self.source, model_id=self.model_id,
|
||||
dimensions=self.dimensions, fallback_reason=None,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def runtime(monkeypatch):
|
||||
runtime = FakeRuntime()
|
||||
monkeypatch.setattr(routed_vectors, "get_model_routing", lambda: runtime)
|
||||
return runtime
|
||||
|
||||
|
||||
async def seed():
|
||||
apple = await note_service.create_note(
|
||||
title="Apple", markdown="apple orchard", folder=None, tags=[],
|
||||
)
|
||||
banana = await note_service.create_note(
|
||||
title="Banana", markdown="banana grove", folder=None, tags=[],
|
||||
)
|
||||
return apple, banana
|
||||
|
||||
|
||||
@pytest.mark.parametrize("outcome", ["api", "api_failure", "missing_space"])
|
||||
def test_benchmark_reports_actual_embedding_and_fallback(runtime, outcome):
|
||||
from app.benchmarks import service
|
||||
from app.contracts import RAGRunRequest
|
||||
|
||||
async def scenario():
|
||||
apple, banana = await seed()
|
||||
if outcome == "api_failure":
|
||||
runtime.result_override = SimpleNamespace(source="local", fallback_reason="PROVIDER_TIMEOUT")
|
||||
elif outcome == "missing_space":
|
||||
runtime.model_id = "space-without-index"
|
||||
directory = get_settings().benchmark_datasets_path
|
||||
directory.mkdir(parents=True, exist_ok=True)
|
||||
(directory / "routing.json").write_text(json.dumps({
|
||||
"dataset_id": "routing", "kind": "rag", "version": "1",
|
||||
"cases": [{"case_id": "query", "query": "apple", "expected_note_ids": [banana.note_id]}],
|
||||
}), encoding="utf-8")
|
||||
run = await service.create_rag_run(RAGRunRequest(
|
||||
dataset_id="routing", modes=[SearchMode.fts, SearchMode.vector],
|
||||
))
|
||||
await service.wait_for_run(run.run_id)
|
||||
report = service.get_report(run.run_id)
|
||||
assert report.config_snapshot["embedding"]["policy"] == "per_case"
|
||||
fts, vector = report.cases
|
||||
assert fts.embedding == {"source": "not_used"}
|
||||
if outcome == "api":
|
||||
assert vector.embedding["source"] == "api"
|
||||
assert vector.embedding["model_id"] == "space-a"
|
||||
assert vector.embedding["dimensions"] == 3
|
||||
assert vector.retrieved_note_ids[0] == banana.note_id
|
||||
else:
|
||||
assert vector.embedding["source"] == "local"
|
||||
assert vector.embedding["model_id"] == "hash-v1"
|
||||
assert vector.embedding["dimensions"] == 128
|
||||
assert vector.retrieved_note_ids[0] == apple.note_id
|
||||
if outcome == "api_failure":
|
||||
assert vector.embedding["fallback_reason"] == "PROVIDER_TIMEOUT"
|
||||
if outcome == "missing_space":
|
||||
assert vector.embedding["fallback_reason"] == "REMOTE_INDEX_UNAVAILABLE"
|
||||
assert vector.embedding["attempted_space"]["model_id"] == "space-without-index"
|
||||
events = service.get_events(run.run_id)
|
||||
case_events = [e for e in events if e.event.value == "CaseCompleted"]
|
||||
assert case_events[-1].data["embedding"] == vector.embedding
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_embedding_observations_are_isolated_between_concurrent_searches(runtime, monkeypatch):
|
||||
from app.retrieval.provenance import capture_embedding
|
||||
|
||||
async def scenario():
|
||||
await seed()
|
||||
original = runtime.embed
|
||||
|
||||
async def embed(texts):
|
||||
await asyncio.sleep(0)
|
||||
if texts == ["offline"]:
|
||||
raise RuntimeError("private upstream details")
|
||||
return await original(texts)
|
||||
|
||||
monkeypatch.setattr(runtime, "embed", embed)
|
||||
|
||||
async def query(text):
|
||||
with capture_embedding() as observation:
|
||||
await engine.search(SearchRequest(query=text, mode=SearchMode.vector))
|
||||
return observation
|
||||
|
||||
remote, local, another = await asyncio.gather(query("apple"), query("offline"), query("apple"))
|
||||
assert remote["source"] == another["source"] == "api"
|
||||
assert local["source"] == "local"
|
||||
assert local["fallback_reason"] == "REMOTE_EMBEDDING_UNAVAILABLE"
|
||||
assert "fallback_reason" not in remote or remote["fallback_reason"] is None
|
||||
assert "private upstream" not in json.dumps(local)
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("failure", ["cancel", "write"])
|
||||
def test_rebuild_failure_preserves_concurrent_configuration_and_all_indexes(runtime, monkeypatch, failure):
|
||||
from app.container import container
|
||||
from app.contracts import ModelRoutingConfig, ProviderConfig, ProviderType
|
||||
from app.services import task_service
|
||||
|
||||
async def scenario():
|
||||
apple, _ = await seed()
|
||||
task = task_service.create_task(title="before", note_id=apple.note_id)
|
||||
before = {table: [tuple(row) for row in rows(f"SELECT * FROM {table}")]
|
||||
for table in ("notes", "blocks", "blocks_fts", "vec_blocks", "index_meta", "routed_block_vectors")}
|
||||
container.model_routing.update(ModelRoutingConfig())
|
||||
entered, release = asyncio.Event(), asyncio.Event()
|
||||
original_embed = runtime.embed
|
||||
|
||||
async def pending_embed(texts):
|
||||
entered.set()
|
||||
await release.wait()
|
||||
return await original_embed(texts)
|
||||
|
||||
monkeypatch.setattr(runtime, "embed", pending_embed)
|
||||
original_index = index_service.index_note
|
||||
writes = 0
|
||||
|
||||
async def fail_write(parsed, **kwargs):
|
||||
nonlocal writes
|
||||
await original_index(parsed, **kwargs)
|
||||
writes += 1
|
||||
if writes == 2:
|
||||
raise RuntimeError("injected write failure")
|
||||
|
||||
if failure == "write":
|
||||
monkeypatch.setattr(index_service, "index_note", fail_write)
|
||||
rebuilding = asyncio.create_task(index_service.rebuild(IndexRebuildRequest()))
|
||||
await asyncio.wait_for(entered.wait(), timeout=5)
|
||||
saved = container.model_routing.update(container.model_routing.configuration())
|
||||
config = ProviderConfig(provider_id="concurrent", provider_type=ProviderType.openai_compatible,
|
||||
name="saved during rebuild", base_url="https://unused.invalid/v1")
|
||||
container.providers.register(config, container.provider_factory.build(config))
|
||||
task_service.update_task(task.task_id, {"title": "saved during rebuild"})
|
||||
# Preparation keeps the old searchable index intact while API I/O is pending.
|
||||
assert repository.stats()["notes"] == 2
|
||||
if failure == "cancel":
|
||||
rebuilding.cancel()
|
||||
expected = asyncio.CancelledError
|
||||
else:
|
||||
release.set()
|
||||
expected = RuntimeError
|
||||
with pytest.raises(expected):
|
||||
await rebuilding
|
||||
assert container.model_routing.configuration().version == saved.config.version
|
||||
assert rows("SELECT provider_id FROM provider_configs")[-1][0] == "concurrent"
|
||||
restored = task_service.get_task(task.task_id)
|
||||
assert restored.title == "saved during rebuild"
|
||||
assert restored.note_id == apple.note_id
|
||||
for table, values in before.items():
|
||||
assert [tuple(row) for row in rows(f"SELECT * FROM {table}")] == values
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def local_engine():
|
||||
return RetrievalEngine(HashEmbeddingProvider(), LexicalReranker(), SqliteVecStore())
|
||||
|
||||
|
||||
def request(mode=SearchMode.vector):
|
||||
return SearchRequest(query="apple", mode=mode, limit=10)
|
||||
|
||||
|
||||
def rows(sql, parameters=()):
|
||||
conn = connect()
|
||||
try:
|
||||
return conn.execute(sql, parameters).fetchall()
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def test_api_index_and_query_use_matching_space_and_keep_local_metadata(runtime):
|
||||
async def scenario():
|
||||
apple, banana = await seed()
|
||||
result = await engine.search(request())
|
||||
assert result.items[0].note_id == banana.note_id
|
||||
baseline = await local_engine().search(request())
|
||||
assert baseline.items[0].note_id == apple.note_id
|
||||
assert rows("SELECT DISTINCT space_id, dimensions FROM routed_block_vectors")[0][:] == ("space-a", 3)
|
||||
assert rows("SELECT COUNT(*) FROM routed_block_vectors")[0][0] == len(apple.blocks) + len(banana.blocks)
|
||||
meta = repository.get_index_meta()
|
||||
assert meta["embedding_model"] == "hash-v1"
|
||||
assert meta["embedding_dim"] == "128"
|
||||
assert len(runtime.calls) == 3
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("failure", ["exception", "local", "missing", "dimension", "corrupt"])
|
||||
def test_query_falls_back_to_exact_local_results(runtime, failure):
|
||||
async def scenario():
|
||||
await seed()
|
||||
if failure == "exception":
|
||||
runtime.error = RuntimeError("offline")
|
||||
elif failure == "local":
|
||||
runtime.source = "local"
|
||||
elif failure == "missing":
|
||||
rows("DELETE FROM routed_block_vectors WHERE block_id = (SELECT MIN(block_id) FROM blocks)")
|
||||
elif failure == "dimension":
|
||||
runtime.dimensions = 4
|
||||
else:
|
||||
rows("UPDATE routed_block_vectors SET vector = ?", ("[0, 0, 0]",))
|
||||
actual = await engine.search(request())
|
||||
baseline = await local_engine().search(request())
|
||||
assert actual == baseline
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_same_dimension_model_switch_never_combines_partial_spaces(runtime):
|
||||
async def scenario():
|
||||
apple, banana = await seed()
|
||||
baseline = await local_engine().search(request())
|
||||
runtime.model_id = "space-b"
|
||||
assert await engine.search(request()) == baseline
|
||||
await note_service.update_note(apple.note_id, markdown="apple orchard")
|
||||
assert {row[0] for row in rows("SELECT DISTINCT space_id FROM routed_block_vectors")} == {"space-a", "space-b"}
|
||||
assert await routed_vectors.search_remote("apple", top_k=10) is None
|
||||
assert await engine.search(request()) == baseline
|
||||
runtime.model_id = "space-a"
|
||||
assert await engine.search(request()) == baseline
|
||||
runtime.model_id = "space-b"
|
||||
await note_service.update_note(banana.note_id, markdown="banana grove")
|
||||
hits = await routed_vectors.search_remote("apple", top_k=10)
|
||||
assert hits is not None and hits[0].id == banana.blocks[0].block_id
|
||||
assert (await engine.search(request())).items[0].note_id == banana.note_id
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_complete_spaces_coexist_but_only_requested_space_is_ranked(runtime):
|
||||
async def scenario():
|
||||
apple, banana = await seed()
|
||||
conn = connect()
|
||||
try:
|
||||
with transaction(conn):
|
||||
routed_vectors.store_remote(
|
||||
conn, [apple.blocks[0].block_id, banana.blocks[0].block_id],
|
||||
routed_vectors.RemoteEmbeddings("space-b", 3, [[1, 0, 0], [0, 1, 0]]),
|
||||
)
|
||||
finally:
|
||||
conn.close()
|
||||
assert (await engine.search(request())).items[0].note_id == banana.note_id
|
||||
runtime.model_id = "space-b"
|
||||
result = await engine.search(request())
|
||||
assert len(result.items) == 2
|
||||
assert result.items[0].note_id == apple.note_id
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_failed_note_embedding_preserves_save_and_forces_coverage_fallback(runtime):
|
||||
async def scenario():
|
||||
apple, banana = await seed()
|
||||
runtime.error = RuntimeError("offline")
|
||||
await note_service.update_note(banana.note_id, markdown="banana changed")
|
||||
assert (await note_service.get_note(banana.note_id)).markdown == "banana changed"
|
||||
assert rows("SELECT COUNT(*) FROM routed_block_vectors")[0][0] == len(apple.blocks)
|
||||
runtime.error = None
|
||||
assert await engine.search(request()) == await local_engine().search(request())
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("vectors, dimensions, space", [
|
||||
([], 3, "space-a"),
|
||||
([[1, 0]], 3, "space-a"),
|
||||
([[0, 0, 0]], 3, "space-a"),
|
||||
([[float("nan"), 0, 0]], 3, "space-a"),
|
||||
([[float("inf"), 0, 0]], 3, "space-a"),
|
||||
([[True, 0, 0]], 3, "space-a"),
|
||||
([[1, 0, 0]], 0, "space-a"),
|
||||
([[1, 0, 0]], 3, "hash-v1"),
|
||||
])
|
||||
def test_invalid_remote_batch_does_not_break_note_saving(runtime, vectors, dimensions, space):
|
||||
runtime.result_override = SimpleNamespace(
|
||||
source="api", vectors=vectors, dimensions=dimensions, model_id=space,
|
||||
)
|
||||
|
||||
async def scenario():
|
||||
note = await note_service.create_note(title="Apple", markdown="apple orchard", folder=None, tags=[])
|
||||
assert (await local_engine().search(request())).items[0].note_id == note.note_id
|
||||
assert await routed_vectors.search_remote("apple", top_k=10) is None
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_remote_storage_failure_rolls_back_batch_but_keeps_local_index(runtime):
|
||||
async def scenario():
|
||||
await seed()
|
||||
rows("""CREATE TRIGGER reject_remote_vector BEFORE INSERT ON routed_block_vectors
|
||||
WHEN (SELECT content FROM blocks WHERE block_id = NEW.block_id) = 'second'
|
||||
BEGIN SELECT RAISE(ABORT, 'simulated storage failure'); END""")
|
||||
note = await note_service.create_note(
|
||||
title="Multi", markdown="first\n\nsecond", folder=None, tags=[],
|
||||
)
|
||||
assert len(note.blocks) == 2
|
||||
assert rows(
|
||||
"SELECT COUNT(*) FROM routed_block_vectors r JOIN blocks b USING(block_id) WHERE b.note_id = ?",
|
||||
(note.note_id,),
|
||||
)[0][0] == 0
|
||||
assert rows("SELECT COUNT(*) FROM vec_blocks")[0][0] == rows("SELECT COUNT(*) FROM blocks")[0][0]
|
||||
assert (get_settings().vault_path / note.file_path).exists()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_rebuild_and_delete_clear_old_remote_rows_through_foreign_keys(runtime):
|
||||
async def scenario():
|
||||
apple, _ = await seed()
|
||||
await note_service.delete_note(apple.note_id)
|
||||
assert rows("SELECT COUNT(*) FROM routed_block_vectors")[0][0] == 1
|
||||
runtime.source = "local"
|
||||
job = await index_service.rebuild(IndexRebuildRequest())
|
||||
assert job.status == "completed"
|
||||
assert rows("SELECT COUNT(*) FROM routed_block_vectors")[0][0] == 0
|
||||
assert rows("SELECT COUNT(*) FROM vec_blocks")[0][0] == 1
|
||||
runtime.source = "api"
|
||||
runtime.model_id = "space-b"
|
||||
await index_service.rebuild(IndexRebuildRequest())
|
||||
assert [row[0] for row in rows("SELECT space_id FROM routed_block_vectors")] == ["space-b"]
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("operation", ["save", "query", "rebuild"])
|
||||
def test_cancellation_propagates_and_mutations_roll_back(runtime, operation):
|
||||
async def scenario():
|
||||
apple, _ = await seed()
|
||||
before = [tuple(row) for row in rows("SELECT * FROM routed_block_vectors ORDER BY block_id")]
|
||||
runtime.error = asyncio.CancelledError()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
if operation == "query":
|
||||
await engine.search(request())
|
||||
elif operation == "rebuild":
|
||||
await index_service.rebuild(IndexRebuildRequest())
|
||||
else:
|
||||
await note_service.update_note(apple.note_id, markdown="changed")
|
||||
assert (await note_service.get_note(apple.note_id)).markdown == "apple orchard"
|
||||
assert [tuple(row) for row in rows("SELECT * FROM routed_block_vectors ORDER BY block_id")] == before
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("injected", ["embedding", "vector_store", "constructor"])
|
||||
def test_injected_engine_dependencies_are_respected(runtime, monkeypatch, injected):
|
||||
async def scenario():
|
||||
apple, _ = await seed()
|
||||
target = engine
|
||||
if injected == "constructor":
|
||||
target = local_engine()
|
||||
elif injected == "embedding":
|
||||
monkeypatch.setattr(engine, "embedding", HashEmbeddingProvider())
|
||||
else:
|
||||
class FakeStore:
|
||||
async def search(self, vector, *, top_k):
|
||||
assert len(vector) == 128
|
||||
return [VectorHit(id=apple.blocks[0].block_id, score=1.0)]
|
||||
|
||||
monkeypatch.setattr(engine, "vector_store", FakeStore())
|
||||
runtime.calls.clear()
|
||||
assert (await target.search(request())).items[0].note_id == apple.note_id
|
||||
assert runtime.calls == []
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_fts_skips_routing_and_hybrid_uses_routed_vector_channel(runtime, monkeypatch):
|
||||
async def scenario():
|
||||
_, banana = await seed()
|
||||
runtime.calls.clear()
|
||||
await engine.search(request(SearchMode.fts))
|
||||
assert runtime.calls == []
|
||||
# Empty lexical channel isolates the vector contribution to hybrid fusion.
|
||||
monkeypatch.setattr(repository, "fts_search", lambda *_: [])
|
||||
|
||||
class PreserveOrder:
|
||||
async def rerank(self, query, candidates):
|
||||
return sorted(candidates, key=lambda candidate: -candidate.score)
|
||||
|
||||
monkeypatch.setattr(engine, "reranker", PreserveOrder())
|
||||
result = await engine.search(request(SearchMode.hybrid))
|
||||
assert result.items[0].note_id == banana.note_id
|
||||
assert runtime.calls == [["apple"]]
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_arbitrary_dimensions_and_extreme_finite_values(runtime):
|
||||
dimensions = 257
|
||||
runtime.result_override = SimpleNamespace(
|
||||
source="api", model_id="space-wide", dimensions=dimensions,
|
||||
vectors=[[1e308, 1e308] + [0.0] * (dimensions - 2)],
|
||||
)
|
||||
|
||||
async def scenario():
|
||||
note = await note_service.create_note(title="Apple", markdown="apple orchard", folder=None, tags=[])
|
||||
hits = await routed_vectors.search_remote("apple", top_k=1)
|
||||
assert hits is not None and hits[0].id == note.blocks[0].block_id
|
||||
assert hits[0].score == pytest.approx(1.0)
|
||||
vector = json.loads(rows("SELECT vector FROM routed_block_vectors")[0][0])
|
||||
assert len(vector) == dimensions
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_missing_runtime_uses_unchanged_local_retrieval(runtime, monkeypatch):
|
||||
monkeypatch.setattr(routed_vectors, "get_model_routing", lambda: None)
|
||||
|
||||
async def scenario():
|
||||
await seed()
|
||||
assert await engine.search(request()) == await local_engine().search(request())
|
||||
assert runtime.calls == []
|
||||
|
||||
asyncio.run(scenario())
|
||||
@@ -30,8 +30,10 @@
|
||||
|
||||
- [AI Core 与 Agent Core 开发说明](development/AI-Core与Agent-Core开发说明.md)
|
||||
- [Knowledge 与 Retrieval Core 开发说明](development/Knowledge与Retrieval-Core开发说明.md)
|
||||
- [Benchmark 开发说明](development/Benchmark开发说明.md)
|
||||
- [模型提供商与模型发现开发说明](development/模型提供商与模型发现开发说明.md)
|
||||
- [MCP Bridge 与 Plugin Host 开发说明](development/MCP-Bridge与Plugin-Host开发说明.md)
|
||||
- [独立 MCP Server 配置中心开发说明](development/独立MCP-Server配置中心开发说明.md)
|
||||
- [Plugin Command 与 Settings 开发说明](development/Plugin-Command与Settings开发说明.md)
|
||||
- [前端壳子与接口层开发说明](development/前端壳子与接口层开发说明.md)
|
||||
- [前端写作体验优化开发说明](development/前端写作体验优化开发说明.md)
|
||||
@@ -49,6 +51,7 @@
|
||||
- [后端全面审阅问题与修复复盘](retrospectives/后端全面审阅问题与修复复盘.md)
|
||||
- [Agent Core 第二阶段问题与修复复盘](retrospectives/Agent-Core第二阶段问题与修复复盘.md)
|
||||
- [Knowledge 与 Retrieval Core 问题与修复复盘](retrospectives/Knowledge与Retrieval-Core问题与修复复盘.md)
|
||||
- [Plugin Command 与 Settings 问题与修复复盘](retrospectives/Plugin-Command与Settings问题与修复复盘.md)
|
||||
- [前端合并审阅问题与修复复盘](retrospectives/前端合并审阅问题与修复复盘.md)
|
||||
|
||||
## 推荐阅读顺序
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
> 适用范围:桌面客户端、本地知识库、RAG、Agent、Skill、多模型接入、多模态处理与可选云同步
|
||||
> 目标读者:前端、Rust 桌面端、Python AI Core、算法、测试与后续接手项目的开发成员
|
||||
|
||||
> 实施状态更新:2026-09-02。本文同时包含目标架构、当前实现和第二阶段接口基线。第一阶段已完成 Vue Web 联调前端、FastAPI、Knowledge/Retrieval、Agent/Tool/Permission、Skill/Plugin 声明式运行时、Mock/OpenAI-Compatible/Ollama Provider、DeepSeek/OpenAI 预设、模型发现及开发阶段 Fernet 凭据存储。Web Workspace 已通过 FastAPI 接入后端配置的真实单 Vault;第二阶段 Agent Trace 持久化、分页快照、可恢复 SSE、stdio MCP Bridge、隔离 Plugin Host、Plugin Command 与 Plugin Settings/Secret Contract 已完成。后续继续接入真实音频处理、Provider 协议增强、Benchmark、文档导出、主题包、Trace 可视化、Mermaid 和函数图像。Tauri/Rust Host、Stronghold、原生多 Vault 文件系统和 Sync Server 仍未实现。
|
||||
> 实施状态更新:2026-09-04。本文同时包含目标架构、当前实现和第二阶段接口基线。第一阶段已完成 Vue Web 联调前端、FastAPI、Knowledge/Retrieval、Agent/Tool/Permission、Skill/Plugin 声明式运行时、Mock/OpenAI-Compatible/Ollama Provider、DeepSeek/OpenAI 预设、模型发现及开发阶段 Fernet 凭据存储。Web Workspace 已通过 FastAPI 接入后端配置的真实单 Vault;第二阶段 Agent Trace 持久化、分页快照、可恢复 SSE、stdio MCP Bridge、隔离 Plugin Host、Plugin Command 与 Plugin Settings/Secret Contract 已完成。阶段 E 已完成 Responses/Anthropic 协议、国内 logo 预设、Provider 配置恢复和 Embedding/转写/声纹 API 路由;本地语音模型仍为阶段 F 接口预留。RAG Benchmark 检索评测(Dataset 加载、异步运行、SSE 进度、指标聚合与报告)已完成,Agent Benchmark 暂缓。后续继续接入真实音频处理、文档导出、主题包、Trace 可视化、Mermaid 和函数图像。Tauri/Rust Host、Stronghold、原生多 Vault 文件系统和 Sync Server 仍未实现。
|
||||
|
||||
---
|
||||
|
||||
@@ -1050,7 +1050,7 @@ Python 包形式的 MCP Server 推荐使用固定版本的 `uvx --isolated --fro
|
||||
|
||||
MCP Bridge 用于接入具有 MCP Server 接口的插件或外部工具服务。
|
||||
|
||||
当前已实现本地 stdio 首版:Plugin Runtime 在授权后的启用阶段启动独立 Server 进程,完成 `initialize`、capability negotiation、分页 `tools/list`、`tools/call`、取消、超时、异常退出和 Host Restart。实现接受 `2025-11-25`、`2025-06-18`、`2025-03-26` 与 `2024-11-05` 协议版本;Streamable HTTP、Resource、Prompt、Sampling 与操作系统级沙箱仍属于后续范围。
|
||||
当前 Plugin Runtime 已实现本地 stdio Host:在授权后的启用阶段启动独立 Server 进程,完成 `initialize`、capability negotiation、分页 `tools/list`、`tools/call`、取消、超时、异常退出和 Host Restart。独立 MCP Server Registry 另行支持 stdio、Streamable HTTP 与旧 HTTP+SSE 兼容,包括 Session、协议 Header、认证 Header Secret、测试门禁和 Tool 动态映射。实现接受 `2025-11-25`、`2025-06-18`、`2025-03-26` 与 `2024-11-05` 协议版本;Plugin Manifest 的 Streamable HTTP、Resource、Prompt、Sampling 与操作系统级沙箱仍属于后续范围。
|
||||
|
||||
MCP Tool 进入系统后的调用路径为:
|
||||
|
||||
@@ -2100,9 +2100,13 @@ MRR
|
||||
Citation Hit Rate
|
||||
P50 Latency
|
||||
P95 Latency
|
||||
total_cases
|
||||
successful_cases
|
||||
failed_cases
|
||||
failure_rate
|
||||
```
|
||||
|
||||
Benchmark 参数、Embedding 模型、Reranker、数据集版本和运行环境需要一起记录,保证不同实验结果可以复现。
|
||||
失败样本按零分计入质量指标分母,报告同时输出样本构成字段标明实际分母。Benchmark 参数、Embedding 模型、Reranker、数据集版本和运行环境需要一起记录,保证不同实验结果可以复现。
|
||||
|
||||
### 20.3 Agent Benchmark
|
||||
|
||||
@@ -2329,7 +2333,7 @@ Markdown Workspace
|
||||
|
||||
第一阶段 Plugin Runtime 已完成安装、启用、停用、权限和声明式 Tool 注册,建立 Skill 调用 Plugin Tool 的基础链路。Command、Settings 和 MCP 执行不计入第一阶段完成项。
|
||||
|
||||
截至 2026-09-02,上述第一阶段后端链路和 Web 联调前端均已完成;第二阶段前置的 Workspace 去 Mock 联调、Agent Trace 持久化/恢复接口、stdio MCP Bridge / Plugin Host 以及 Plugin Command/Settings 后端 Contract 也已完成。当前验证基线为后端 126 项测试、前端 29 项测试、TypeScript 类型检查及生产构建通过。向量链路当前使用 `HashEmbeddingProvider` 验证工程正确性,真实 Embedding 召回质量不属于该测试结论。
|
||||
截至 2026-09-03,上述第一阶段后端链路和 Web 联调前端均已完成;第二阶段的 Workspace 去 Mock 联调、Agent Trace 持久化/恢复接口、stdio MCP Bridge / Plugin Host,以及 Plugin Command/Settings 前后端闭环也已完成。Plugin 详情页现已提供 Host 状态、重启、动态设置、Secret 管理和命令执行,全局命令面板可加载 Plugin Command。当前验证基线为后端 136 项测试、前端 32 项测试、TypeScript 类型检查及生产构建通过。向量链路当前使用 `HashEmbeddingProvider` 验证工程正确性,真实 Embedding 召回质量不属于该测试结论。
|
||||
|
||||
第二阶段在既有 Contract 上接入:
|
||||
|
||||
@@ -2361,7 +2365,7 @@ Frontend Extension
|
||||
└── Plugin Settings UI
|
||||
```
|
||||
|
||||
上述列表描述第二阶段技术范围,其中 stdio MCP Bridge 已实现,其余能力以各自开发说明的状态为准。每项功能必须继续经过现有 Service、Contract、Permission 和 Adapter 边界,不因 Demo 需要在 Vue 组件、Router 或 Agent Runtime 中直接绑定第三方协议。
|
||||
上述列表描述第二阶段技术范围,其中 stdio MCP Bridge、Plugin Command Contribution 和 Plugin Settings Contribution 后端 Contract 已实现,其余能力以各自开发说明的状态为准。每项功能必须继续经过现有 Service、Contract、Permission 和 Adapter 边界,不因 Demo 需要在 Vue 组件、Router 或 Agent Runtime 中直接绑定第三方协议。
|
||||
|
||||
第三阶段处理:
|
||||
|
||||
@@ -2415,9 +2419,9 @@ Sync Server 按独立服务开发和部署,不进入桌面客户端核心启
|
||||
|
||||
目标桌面端采用 Tauri 2、Rust、Vue 3 和 TypeScript;当前可运行形态是 Vue/Vite Web 前端加 FastAPI。用户笔记以 Markdown 和 Assets 保存在本地 Vault,SQLite 已管理笔记元数据、全文索引、向量索引、任务及 Agent Trace;Provider/Extension Registry 当前仍为内存实现。
|
||||
|
||||
Python AI Core 未来作为 Tauri Sidecar 运行,当前由开发命令独立启动,FastAPI 提供本地接口。Knowledge Core 管理笔记结构;Retrieval Core 当前通过 FTS5、`HashEmbeddingProvider`、sqlite-vec、RRF 和轻量 Reranker 跑通混合检索,真实 Embedding 与正式 Benchmark 在第二阶段接入;Agent Runtime 使用 Tool Registry 操作知识库和任务,并将扩展 Agent Trace Contract 供可视化和 Benchmark 共用;Skill Runtime 将提示词、工具、权限和检索参数组装为可复用 Agent 配置。
|
||||
Python AI Core 未来作为 Tauri Sidecar 运行,当前由开发命令独立启动,FastAPI 提供本地接口。Knowledge Core 管理笔记结构;Retrieval Core 当前通过 FTS5、`HashEmbeddingProvider`、sqlite-vec、RRF 和轻量 Reranker 跑通混合检索,真实 Embedding 与正式 Benchmark 仍待第二阶段后续接入;Agent Runtime 使用 Tool Registry 操作知识库和任务,并已持久化可供前端可视化与 Benchmark 共用的 Agent Trace Contract;Skill Runtime 将提示词、工具、权限和检索参数组装为可复用 Agent 配置。
|
||||
|
||||
当前 Plugin Runtime 支持 Manifest、生命周期和声明式白名单 Tool Contribution,并已通过 stdio MCP Bridge 接入独立进程 Tool、Host 状态与重启接口;Command 与 Settings Contribution 尚待后续阶段实现。Provider Adapter 当前实现 Mock、OpenAI Chat/OpenAI-Compatible 与 Ollama,第二阶段按统一行为测试完善 OpenAI Responses、Anthropic Messages 等协议。多模态目标方案使用 faster-whisper、pyannote.audio 和可选 emotion2vec;当前只读取 Host 预生成 transcript。
|
||||
当前 Plugin Runtime 支持 Manifest、生命周期、声明式白名单 Tool Contribution、Plugin Command 与 Plugin Settings/Secret,并已通过 stdio MCP Bridge 接入独立进程 Tool、专用 MCP Command Target、Host 状态与重启接口。Provider Adapter 当前实现 Mock、OpenAI Chat/OpenAI-Compatible 与 Ollama,OpenAI Responses、Anthropic Messages 等协议仍待第二阶段后续完善。多模态目标方案使用 faster-whisper、pyannote.audio 和可选 emotion2vec;当前只读取 Host 预生成 transcript。
|
||||
|
||||
第二阶段内容输出以 Document AST、Exporter Adapter、Mermaid Renderer 和 Function Plot Renderer 为共同边界,支持 HTML、PDF、DOCX 与静态图导出。Theme Package 使用 Manifest、Design Token 和受限 CSS 实现本地导入;联网主题市场不属于本阶段核心依赖。API Key 在 Web 联调期由 Fernet 开发存储加密保存,桌面版迁移到 Tauri Stronghold。多设备同步的目标方案为独立、可自托管的 Sync Server,目前尚未实现;本地核心功能不依赖 Sync Server。
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# 第一阶段分工表
|
||||
|
||||
> 状态更新:2026-08-30。本文保留原始职责划分,同时记录当前交付状态。第一阶段后端目标已完成,前端 Web 联调页面已完成;尚未纳入本阶段完成项的是 Tauri Host、Stronghold、真实桌面文件系统、独立 MCP Plugin Host、真实音频模型和 Sync Server。
|
||||
> 状态更新:2026-09-02。本文保留第一阶段原始职责划分和交付口径。第一阶段后端目标与前端 Web 联调页面均已完成;第二阶段此后又完成 Agent Trace 持久化与 SSE 恢复、独立 stdio MCP Plugin Host 和 Plugin Command/Settings 后端 Contract。Tauri Host、Stronghold、原生多 Vault 文件系统、真实音频模型和 Sync Server 仍未实现。
|
||||
|
||||
## 当前交付状态
|
||||
|
||||
|
||||
@@ -6,9 +6,9 @@
|
||||
|
||||
第二阶段继续保持第一阶段的模块 ownership:
|
||||
|
||||
- 范涵宇:Agent Core、Extension Core、Model Core、Multimodal、整体架构与代码审阅。
|
||||
- 范涵宇:Agent Core、Extension Core、Model Core、Multimodal、Plugin Command / Settings UI、整体架构与代码审阅。
|
||||
- 杨星萱:Knowledge Core、Retrieval Core、Benchmark、文档导出、函数图像绘制。
|
||||
- 吉海燕:Frontend、Theme、Agent Trace、Plugin UI Contribution、Mermaid 渲染。
|
||||
- 吉海燕:Frontend、Theme、Agent Trace、Mermaid 渲染。
|
||||
|
||||
---
|
||||
|
||||
@@ -16,9 +16,9 @@
|
||||
|
||||
| 成员 | 主要负责方向 | 第二阶段任务 | 配合事项 |
|
||||
| --- | --- | --- | --- |
|
||||
| 范涵宇 | Agent Core / Extension Core / Model Core / Multimodal / 总体架构 | faster-whisper、pyannote.audio、MCP Bridge、Plugin Command Contribution 后端、Plugin Settings Contribution 后端、更多 Provider、整体集成、代码审阅与统筹 | 与吉海燕联调 Plugin Command / Settings 前端;为杨星萱的 Agent Benchmark 提供 Agent Trace、Tool Call 等测试接口 |
|
||||
| 范涵宇 | Agent Core / Extension Core / Model Core / Multimodal / 总体架构 | faster-whisper、pyannote.audio、MCP Bridge、Plugin Command / Settings 前后端闭环、更多 Provider、整体集成、代码审阅与统筹 | 为杨星萱的 Agent Benchmark 提供 Agent Trace、Tool Call 等测试接口 |
|
||||
| 杨星萱 | Knowledge Core / Retrieval Core / Benchmark / Export / 数学内容渲染 | RAG Benchmark、Agent Benchmark 基础设施、Markdown → HTML / PDF / DOCX、函数图像绘制与渲染支持、Retrieval 调优 | 与范涵宇确认 Agent Benchmark 事件和测试数据结构;与吉海燕联调函数图像在编辑器和预览区中的展示 |
|
||||
| 吉海燕 | Frontend / Theme / Visualization | Theme Import、Theme Manifest、社区主题格式、Agent Trace 可视化、Plugin Command / Settings 前端、Mermaid 渲染支持 | 与范涵宇联调 Plugin Contribution Contract 与 AgentEvent;与杨星萱联调函数图像及导出预览 |
|
||||
| 吉海燕 | Frontend / Theme / Visualization | Theme Import、Theme Manifest、社区主题格式、Agent Trace 可视化、Mermaid 渲染支持 | 与范涵宇联调 AgentEvent;与杨星萱联调函数图像及导出预览 |
|
||||
|
||||
---
|
||||
|
||||
@@ -384,7 +384,7 @@ Metadata Filter
|
||||
|
||||
---
|
||||
|
||||
## 五、吉海燕
|
||||
## 五、前端展示层(吉海燕;5.4—5.5 由范涵宇负责)
|
||||
|
||||
### 5.1 Theme Import
|
||||
|
||||
@@ -488,7 +488,7 @@ Error
|
||||
- 用户取消;
|
||||
- Citation 跳转。
|
||||
|
||||
### 5.4 Plugin Command 前端
|
||||
### 5.4 Plugin Command 前端(范涵宇)
|
||||
|
||||
负责 Command Contribution 在前端呈现。
|
||||
|
||||
@@ -502,7 +502,7 @@ Toolbar Action
|
||||
|
||||
前端使用 Plugin Contribution Contract,不直接解析插件后端 Manifest。
|
||||
|
||||
### 5.5 Plugin Settings 前端
|
||||
### 5.5 Plugin Settings 前端(范涵宇)
|
||||
|
||||
根据范涵宇提供的 Plugin Settings Schema 动态生成设置表单。
|
||||
|
||||
@@ -572,9 +572,9 @@ Mermaid 渲染需要与 Theme Design Token 联动。
|
||||
| --- | --- | --- |
|
||||
| MCP Bridge → Agent Tool | 范涵宇 | 杨星萱 |
|
||||
| Plugin Command Runtime | 范涵宇 | 吉海燕 |
|
||||
| Plugin Command UI | 吉海燕 | 范涵宇 |
|
||||
| Plugin Command UI | 范涵宇 | 吉海燕 |
|
||||
| Plugin Settings Runtime | 范涵宇 | 吉海燕 |
|
||||
| Plugin Settings UI | 吉海燕 | 范涵宇 |
|
||||
| Plugin Settings UI | 范涵宇 | 吉海燕 |
|
||||
| Agent Benchmark Framework | 杨星萱 | 范涵宇 |
|
||||
| Agent Trace Event Contract | 范涵宇 | 吉海燕、杨星萱 |
|
||||
| Agent Trace Visualization | 吉海燕 | 范涵宇 |
|
||||
@@ -750,6 +750,8 @@ Markdown
|
||||
- [x] Agent 能调用 MCP Tool;
|
||||
- [x] Plugin Command Contribution 后端可注册;
|
||||
- [x] Plugin Settings Contribution 后端可解析;
|
||||
- [x] Plugin Command 可以显示并从前端执行;
|
||||
- [x] Plugin Settings 可以动态生成设置项并独立提交 Secret;
|
||||
- [ ] Provider Adapter 的 Streaming / Tool Calling / Error Mapping 稳定;
|
||||
- [ ] 完成跨模块接口审阅和第二阶段集成。
|
||||
|
||||
@@ -773,8 +775,6 @@ Markdown
|
||||
- [ ] Theme 可以启用、停用和卸载;
|
||||
- [ ] Agent Trace 可以展示完整 Tool Call 顺序;
|
||||
- [ ] Trace Node 可以查看参数、结果、耗时和错误;
|
||||
- [ ] Plugin Command 可以显示在前端;
|
||||
- [ ] Plugin Settings 可以动态生成设置项;
|
||||
- [ ] Markdown Mermaid Code Block 可以渲染;
|
||||
- [ ] Mermaid 支持主题切换;
|
||||
- [ ] Mermaid 渲染错误可以明确展示;
|
||||
@@ -791,6 +791,8 @@ Markdown
|
||||
├── MCP Bridge
|
||||
├── Plugin Command Runtime
|
||||
├── Plugin Settings Runtime
|
||||
├── Plugin Command UI
|
||||
├── Plugin Settings UI
|
||||
├── Provider Adapter
|
||||
├── Code Review
|
||||
└── Integration / Coordination
|
||||
@@ -812,8 +814,6 @@ Markdown
|
||||
├── Theme Import
|
||||
├── Theme Community Format
|
||||
├── Agent Trace Visualization
|
||||
├── Plugin Command UI
|
||||
├── Plugin Settings UI
|
||||
└── Mermaid
|
||||
├── Markdown Rendering
|
||||
├── Theme Adaptation
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# 后端接口契约(开发版)
|
||||
|
||||
> 更新日期:2026-09-01。本文档记录当前前后端联调使用的已实现接口;机器可读字段、校验规则和响应模型以 FastAPI 运行时生成的 OpenAPI 为准。第二阶段尚未实现的规划接口见 `第二阶段接口契约-开发版.md`,不要将规划路径视为当前服务能力。
|
||||
> 更新日期:2026-09-02。本文档记录当前前后端联调使用的已实现接口;机器可读字段、校验规则和响应模型以 FastAPI 运行时生成的 OpenAPI 为准。第二阶段尚未实现的规划接口见 `第二阶段接口契约-开发版.md`,不要将规划路径视为当前服务能力。
|
||||
|
||||
## 契约入口
|
||||
|
||||
@@ -176,11 +176,11 @@ RunCancelled
|
||||
|
||||
## 当前实现状态
|
||||
|
||||
更新至 2026-09-02:后端 126 项回归测试通过;第二阶段 Plugin Command 与 Plugin Settings/Secret 接口已实现,详细 DTO 和边界见《第二阶段接口契约-开发版》第 7 节。
|
||||
更新至 2026-09-02:后端 136 项回归测试通过;第二阶段 Plugin Command 与 Plugin Settings/Secret 接口已实现,详细 DTO 和边界见《第二阶段接口契约-开发版》第 7 节。
|
||||
|
||||
- Chat、Agent Run、Agent Events、Tool 列表、Provider 配置生命周期、模型列表和连接测试已经接入 AI Core。
|
||||
- Agent Run/Event 已持久化到 SQLite;SSE 帧携带 sequence `id`,断线后可以回放缺失事件。Trace API 与 Benchmark 共用同一事件事实,并在入库前执行 Secret 脱敏和结果限长。
|
||||
- Provider Adapter 当前包含 Mock、真正增量 SSE 的 OpenAI-Compatible Chat Completions,以及 Ollama JSONL Streaming。
|
||||
- Provider Adapter 当前包含 Mock、增量 SSE 的 OpenAI-Compatible Chat Completions、OpenAI Responses、Anthropic Messages,以及 Ollama JSONL Streaming。阶段 E 增加 `/api/model-routing`、`/api/models/embeddings`、`/api/media/speaker-matches`;具体请求和阶段边界见第二阶段契约 §8.5。
|
||||
- Notes、Search、Index、Skills、Plugins、Tasks 和 Provider 生命周期均已接入业务服务。
|
||||
- Workspace 已接入后端配置的真实 Vault;文件树、笔记读写、文件/目录新建、重命名和删除不再使用前端 Mock Fallback。
|
||||
- Note Move 保留 `note_id`;Citation 的字符偏移统一使用 UTF-16 code unit,供浏览器编辑器直接定位。
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
> 文档状态:接口冻结草案
|
||||
>
|
||||
> 更新日期:2026-09-01
|
||||
> 更新日期:2026-09-03
|
||||
>
|
||||
> 依据:`../architecture/第二阶段团队分工表.md`、`../architecture/AI笔记软件技术栈说明-团队版-v2.3.md`、`后端接口契约-开发版.md`
|
||||
|
||||
@@ -27,7 +27,7 @@
|
||||
|
||||
| 标记 | 含义 |
|
||||
| --- | --- |
|
||||
| 已实现 | 第一阶段接口已经存在,第二阶段保持兼容 |
|
||||
| 已实现 | 接口已经落地并由当前 OpenAPI 与自动化测试覆盖 |
|
||||
| 扩展 | 路径已存在,第二阶段增加字段、事件或行为 |
|
||||
| 计划新增 | 第二阶段需要新增实现 |
|
||||
| 内部 Contract | 不直接暴露 HTTP,由两个模块共同遵守 |
|
||||
@@ -47,6 +47,13 @@
|
||||
| Agent Trace | GET | `/api/agent/runs/{run_id}/trace` | 已实现 | 分页读取可回放 Trace 快照 |
|
||||
| Plugin Host | GET | `/api/plugins/{plugin_id}/host` | 已实现 | 获取 MCP Host 健康状态 |
|
||||
| Plugin Host | POST | `/api/plugins/{plugin_id}/host/restart` | 已实现 | 重启异常 Host 并重新发现 Tool |
|
||||
| MCP Server | GET/POST | `/api/mcp/servers` | 已实现(C.1) | 列出、创建独立 MCP Server 配置 |
|
||||
| MCP Server | GET/PUT/DELETE | `/api/mcp/servers/{server_id}` | 已实现(C.1) | 读取、版本化修改、删除独立配置 |
|
||||
| MCP Server | GET | `/api/mcp/servers/{server_id}/tools` | 已实现(C.1) | 获取映射后的 Tool 摘要 |
|
||||
| MCP Server | POST | `/api/mcp/servers/{server_id}/trust` | 已实现(C.1) | 确认当前连接配置摘要 |
|
||||
| MCP Server | POST | `/api/mcp/servers/{server_id}/test` | 已实现(C.1) | 临时连接、握手、发现工具后关闭 |
|
||||
| MCP Server | POST | `/api/mcp/servers/{server_id}/enable`、`disable` | 已实现(C.1) | 控制连接与动态 Tool 生命周期 |
|
||||
| MCP Server | PUT/DELETE | `/api/mcp/servers/{server_id}/secrets/{key}` | 已实现(C.1) | 按 `kind` 写入或删除加密环境变量/Header |
|
||||
| Plugin Command | GET | `/api/plugin-contributions/commands` | 已实现 | 获取前端可展示的 Command |
|
||||
| Plugin Command | POST | `/api/plugin-contributions/commands/{command_id}/execute` | 已实现 | 受控执行 Command |
|
||||
| Plugin Settings | GET | `/api/plugins/{plugin_id}/settings` | 已实现 | 获取 Schema 与非敏感配置 |
|
||||
@@ -55,9 +62,9 @@
|
||||
| Provider | 现有路径 | `/api/providers/*`、`POST /api/chat` | 扩展 | 补齐协议能力和统一行为 |
|
||||
| Retrieval | GET/POST | `/api/index/status`、`/api/index/rebuild` | 扩展 | 暴露 Embedding 兼容状态并安全重建向量 |
|
||||
| Benchmark | GET | `/api/benchmarks/datasets` | 计划新增 | 枚举受控 Dataset |
|
||||
| Benchmark | POST | `/api/benchmarks/rag/runs` | 计划新增 | 创建 RAG Benchmark |
|
||||
| Benchmark | POST | `/api/benchmarks/agent/runs` | 计划新增 | 创建 Agent Benchmark |
|
||||
| Benchmark | GET | `/api/benchmarks/runs` | 计划新增 | 分页获取 Benchmark Run |
|
||||
| Benchmark | POST | `/api/benchmarks/rag/runs` | 已实现 | 创建 RAG Benchmark |
|
||||
| Benchmark | POST | `/api/benchmarks/agent/runs` | 暂缓 | 创建 Agent Benchmark(依赖 Agent Runtime 完成后交付) |
|
||||
| Benchmark | GET | `/api/benchmarks/runs` | 已实现 | 分页获取 Benchmark Run |
|
||||
| Benchmark | GET/POST | `/api/benchmarks/runs/{run_id}/*` | 计划新增 | 查询、订阅、取消和读取报告 |
|
||||
| Export | POST | `/api/exports` | 计划新增 | 创建 HTML/PDF/DOCX 导出任务 |
|
||||
| Export | GET | `/api/exports` | 计划新增 | 分页获取导出任务 |
|
||||
@@ -647,6 +654,29 @@ MCP_TOOL_SCHEMA_INVALID
|
||||
MCP_TOOL_CALL_FAILED
|
||||
MCP_TOOL_RESULT_TOO_LARGE
|
||||
MCP_TRUST_APPROVAL_REQUIRED
|
||||
MCP_TRUST_DIGEST_STALE
|
||||
MCP_SANDBOX_REQUIRED
|
||||
MCP_TRANSPORT_UNSUPPORTED
|
||||
MCP_SERVER_NOT_FOUND
|
||||
MCP_SERVER_NAME_INVALID
|
||||
MCP_SERVER_ALREADY_ENABLED
|
||||
MCP_SERVER_VERSION_CONFLICT
|
||||
MCP_SERVER_LIMIT_REACHED
|
||||
MCP_REGISTRY_WRITE_FAILED
|
||||
MCP_REGISTRY_INVALID
|
||||
MCP_CONNECTION_TEST_REQUIRED
|
||||
MCP_CONFIG_INVALID
|
||||
MCP_COMMAND_INVALID
|
||||
MCP_URL_INVALID
|
||||
MCP_HEADER_INVALID
|
||||
MCP_HTTP_REQUEST_FAILED
|
||||
MCP_HTTP_RESPONSE_INVALID
|
||||
MCP_SECRET_REQUIRED
|
||||
MCP_SECRET_NOT_DECLARED
|
||||
MCP_SECRET_KIND_INVALID
|
||||
MCP_SECRET_STORE_ERROR
|
||||
MCP_ENVIRONMENT_INVALID
|
||||
MCP_PERMISSION_INVALID
|
||||
PLUGIN_COMMAND_NOT_FOUND
|
||||
PLUGIN_COMMAND_CONFLICT
|
||||
PLUGIN_COMMAND_INVALID
|
||||
@@ -670,10 +700,26 @@ PLUGIN_STORAGE_ERROR
|
||||
CREDENTIAL_NAMESPACE_RESERVED
|
||||
```
|
||||
|
||||
### 7.8 独立 MCP Server Registry(C.1)
|
||||
|
||||
独立 Server 不依附 Plugin Manifest,配置持久化于 `APP_DATA_DIR/mcp/servers.json`。`transport` 支持 `stdio`、`streamable_http` 和兼容旧服务的 `sse`。敏感环境变量与认证 Header 只以 `mcp.*` 引用进入加密凭据存储;读取响应以 `secret_environment`、`secret_headers` 的布尔值表示配置状态,不返回明文。动态工具使用 `mcp.{server_id}.{remote_tool}` 命名空间,来源标记为 `mcp_server`,仍通过统一 Tool Registry、Permission Manager 与 Agent Trace。
|
||||
|
||||
stdio 配置使用 `command`、`args`、`environment` 和 `secret_environment_keys`;HTTP/SSE 配置使用 `url`、`headers` 和 `secret_header_keys`,两组 Transport 字段不可混用。更新请求必须携带当前 `version`,成功后版本递增;过期版本返回 `409 MCP_SERVER_VERSION_CONFLICT`。`GET /tools` 返回 `name`、`remote_name`、`description` 和可选 `permission`。
|
||||
|
||||
创建或编辑配置后,调用方必须向 `/trust` 回传服务端计算的 `command_digest`。后端只接受与当前 Transport、连接参数、环境/Header 及权限完全一致的摘要;配置变化会撤销旧信任和测试结果。只有当前摘要通过 `/test`,才能调用 `/enable`。测试失败也会持久化时间和失败状态。
|
||||
|
||||
Secret 明文变化无法进入摘要,因此 Secret 写入和删除采用更严格规则:若 Server 已启用则先停用并注销 Tool,随后清除 `tested_digest` 和最近测试状态。调用方必须用新 Secret 再次执行 `/test`,不能沿用旧凭据的测试结果。
|
||||
|
||||
Streamable HTTP 支持 Session ID、`MCP-Protocol-Version`、JSON 或 SSE POST 响应、可选 GET 事件流及 `Last-Event-ID`;旧 SSE 按 endpoint 事件确定 POST 地址,并要求与初始 URL 同源。Secret 接口用 `?kind=environment` 或 `?kind=header` 区分类型。HTTP URL 不允许内嵌凭据或 Fragment,配置不得覆盖协议保留 Header。
|
||||
|
||||
stdio 命令始终以 executable 与 args 数组通过 `shell=False` 启动;普通环境变量和加密 Secret 显式注入,不继承 Provider Key、数据库或 Vault 路径。当前 Python Host 仅在 `APP_ENVIRONMENT=development` 时允许启动 stdio;其他环境返回 `403 MCP_SANDBOX_REQUIRED`。远程 HTTP Transport 不创建本机进程,但仍受摘要确认、成功测试、超时、消息限长与 Secret 隔离约束。
|
||||
|
||||
---
|
||||
|
||||
## 8. Provider Adapter 扩展
|
||||
|
||||
> 阶段 E 实施更新(2026-09-04):OpenAI Responses、Anthropic Messages、Chat Completions 与 Ollama Adapter 已接入;国内提供商 logo 预设、独立凭据输入、配置恢复、Embedding / 转写 / 声纹 API 路由已实现。真实本地语音模型仍属于阶段 F。实现细节见 [模型提供商与模型发现开发说明](../development/模型提供商与模型发现开发说明.md)。
|
||||
|
||||
第二阶段不新增平行 Provider CRUD,继续使用第一阶段接口:
|
||||
|
||||
```text
|
||||
@@ -690,7 +736,7 @@ POST /api/chat
|
||||
|
||||
### 8.1 ModelInfo 扩展
|
||||
|
||||
`GET /api/providers/{provider_id}/models` 的 item 增加可选字段:
|
||||
以下为后续计划的可选字段;阶段 E 的 `GET /api/providers/{provider_id}/models` 实际 item 仍只包含 `model`、`display_name`、`capabilities`:
|
||||
|
||||
```json
|
||||
{
|
||||
@@ -730,6 +776,8 @@ Done
|
||||
- 浏览器取消 Fetch 或 SSE 后,服务端必须取消上游 Provider 请求。
|
||||
- 不支持 reasoning 的 Provider 不发送伪造 ThinkingDelta。
|
||||
|
||||
阶段 E 补充:取消或关闭迭代器直接关闭上游连接并传播取消,不向已断开的客户端继续发送 Done。内部带点号、长名称的工具映射为合法的 64 字符以内名称,响应恢复原命名空间,映射在请求内隔离。实际流中断错误码为 `PROVIDER_STREAM_TRUNCATED`;`PROVIDER_INVALID_RESPONSE` 用于无效结构/参数。上面的 `Done.data.status` 适用于真实 HTTP Adapter;开发 Mock 保留原有测试事件。
|
||||
|
||||
### 8.3 Provider 一致性测试 Contract
|
||||
|
||||
每个 Adapter 使用相同 Case 描述:
|
||||
@@ -767,6 +815,25 @@ MODEL_CONTEXT_LENGTH_EXCEEDED
|
||||
|
||||
---
|
||||
|
||||
### 8.5 阶段 E 模型路由接口(已实现)
|
||||
|
||||
| 方法 | 路径 | 契约 |
|
||||
| --- | --- | --- |
|
||||
| GET | `/api/model-routing` | `{config, local_backends}` |
|
||||
| PUT | `/api/model-routing` | 提交 ModelRoutingConfig,返回递增版本配置 |
|
||||
| POST | `/api/models/embeddings` | `{texts: string[]}` → `{vectors, source, model_id, dimensions, fallback_reason}` |
|
||||
| POST | `/api/media/speaker-matches` | `{attachment_id, reference_attachment_id}` → `{score, source, fallback_reason}` |
|
||||
|
||||
`ModelRoutingConfig` 包含 `version`、`embedding`、`transcription`、`speaker_matching`。每个能力为 null 或 `{provider_id, model, endpoint, dimensions?}`。dimensions 仅 Embedding 使用,范围 1–16384;endpoint 是选定 Provider 下不带查询的路径,不能传第二个 URL。PUT 不提交 GET 返回的 local_backends;版本冲突返回 409 `MODEL_ROUTING_VERSION_CONFLICT`。删除仍被引用的 Provider 返回 409 `PROVIDER_IN_USE`。
|
||||
|
||||
本阶段远程能力仅接受 OpenAI Chat / Compatible HTTP 配置,默认路径分别是 `/embeddings`、`/audio/transcriptions`、`/audio/speaker-matches`。最后一个是本项目自定义 multipart 接口,**不是公共 OpenAI 标准协议**:请求 model、file、reference_file,响应有限 0–1 的 score。转写采用 multipart model、file、可选 language,必须返回非空 text。文件来自受控附件目录,限制 25 MiB。
|
||||
|
||||
`TranscriptionJob` 新增可选 `source: api|local|sidecar` 和 `fallback_reason`。保留已有转写 Job 路径;`diarization=true` 返回失败 Job,错误为 `DIARIZATION_NOT_IMPLEMENTED`,不能静默忽略。
|
||||
|
||||
无配置时调用本地接口;有配置时 API 优先,网络/鉴权/限流/结果无效时回退。本地 Embedding 当前为 hash 占位;本地 ASR / 声纹后端尚未安装时返回 `LOCAL_MODEL_NOT_INSTALLED`,而非伪成功。远程 Embedding 独立索引并检查完整覆盖,模型变化后需重建;不与本地向量混算。
|
||||
|
||||
Provider PATCH 支持 provider_type;普通配置持久化到 SQLite,凭据继续独立加密。预设新增 logo_id、description、capabilities,前端图标随应用打包。
|
||||
|
||||
## 9. RAG / Agent Benchmark
|
||||
|
||||
Benchmark Service 同时提供 Python 调用接口和本地 HTTP 接口。CLI、测试和前端报告页调用同一 Service,不各自实现指标。
|
||||
@@ -845,7 +912,7 @@ Dataset 从仓库或受控导入目录注册。API 不接受调用方提交任
|
||||
|
||||
配置快照必须记录 Embedding model ID/version/dimension、Reranker、索引版本、Dataset Hash 和运行环境。
|
||||
|
||||
### 9.5 创建 Agent Benchmark
|
||||
### 9.5 创建 Agent Benchmark(暂缓,未暴露接口)
|
||||
|
||||
`POST /api/benchmarks/agent/runs`
|
||||
|
||||
@@ -879,12 +946,16 @@ RAG 和 Agent 创建接口均返回 `202 BenchmarkRun`:
|
||||
"metrics": null,
|
||||
"config_snapshot": {},
|
||||
"error": null,
|
||||
"error_code": null,
|
||||
"created_at": "2026-08-31T10:30:00Z",
|
||||
"started_at": null,
|
||||
"completed_at": null
|
||||
}
|
||||
```
|
||||
|
||||
`status` 取值:`queued` → `running` → `completed` | `failed` | `cancelled`。失败/取消时 `error` 与
|
||||
`error_code` 只返回项目错误码与安全消息,不暴露第三方堆栈。
|
||||
|
||||
公共接口:
|
||||
|
||||
| 方法 | 路径 | 用途 |
|
||||
@@ -895,6 +966,10 @@ RAG 和 Agent 创建接口均返回 `202 BenchmarkRun`:
|
||||
| POST | `/api/benchmarks/runs/{run_id}/cancel` | 取消运行 |
|
||||
| GET | `/api/benchmarks/runs/{run_id}/report` | 获取结构化完整报告 |
|
||||
|
||||
SSE 事件流(`RunStarted` → `CaseCompleted`* → `RunCompleted` | `RunFailed` | `RunCancelled`):
|
||||
`GET /api/benchmarks/runs/{run_id}/events` 支持 `Last-Event-ID` 与 `?after_sequence=` 游标恢复,
|
||||
`RunCompleted` / `RunFailed` / `RunCancelled` 为终止事件,收到后即断流。
|
||||
|
||||
### 9.7 指标 Contract
|
||||
|
||||
RAG:
|
||||
@@ -907,10 +982,17 @@ RAG:
|
||||
"mrr": 0.81,
|
||||
"citation_hit_rate": 0.89,
|
||||
"p50_latency_ms": 24.5,
|
||||
"p95_latency_ms": 67.3
|
||||
"p95_latency_ms": 67.3,
|
||||
"total_cases": 50,
|
||||
"successful_cases": 48,
|
||||
"failed_cases": 2,
|
||||
"failure_rate": 0.04
|
||||
}
|
||||
```
|
||||
|
||||
失败样本按零分计入质量指标分母,`total_cases` / `successful_cases` / `failed_cases` /
|
||||
`failure_rate` 让报告明确实际分母;延迟仅统计成功样本。
|
||||
|
||||
Agent:
|
||||
|
||||
```json
|
||||
@@ -934,8 +1016,10 @@ BENCHMARK_DATASET_NOT_FOUND
|
||||
BENCHMARK_DATASET_INVALID
|
||||
BENCHMARK_CONFIG_INVALID
|
||||
BENCHMARK_INDEX_INCOMPATIBLE
|
||||
BENCHMARK_CAPACITY_EXCEEDED
|
||||
BENCHMARK_RUN_NOT_FOUND
|
||||
BENCHMARK_RUN_FAILED
|
||||
BENCHMARK_CASE_EVALUATION_FAILED
|
||||
```
|
||||
|
||||
### 9.9 Retrieval Profile 与索引兼容
|
||||
@@ -1483,3 +1567,7 @@ frontend/src/
|
||||
```
|
||||
|
||||
目录调整应按实际代码规模渐进进行。Router 只做参数接收和错误映射,状态机、第三方 SDK 与文件处理继续放在 Service/Adapter 层。
|
||||
|
||||
### Benchmark Embedding 运行归属(阶段 E 集成修复)
|
||||
|
||||
`config_snapshot.local_embedding` 仅表示本地基线;`config_snapshot.embedding` 为 `{ "policy": "per_case", "details": "cases[].embedding" }`。报告与 CaseCompleted 事件的逐样本 `embedding` 包含实际 source(api/local/not_used/unavailable)、model_id、dimensions,以及可选 version、fallback_reason、requested_route、route_version、attempted_space。requested_route 仅含提供商引用、模型、相对端点和维度,不包含 API Key 或凭据引用。FTS 不使用 Embedding,标记 not_used;远程失败或索引不完整回退时记录实际本地模型及原因。
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
> 本文档用于团队开发和模块联调,记录当前已经落地的核心边界与使用方式。
|
||||
|
||||
> 更新日期:2026-09-02。第一阶段 AI Core、Agent Core、Extension Core 和 Model Core 主链路已经完成;第二阶段 Agent Trace 持久化、可恢复 SSE、stdio MCP Bridge、隔离 Plugin Host 以及 Plugin Command/Settings 已落地,后端当前回归基线为 126 项测试通过。
|
||||
> 更新日期:2026-09-02。第一阶段 AI Core、Agent Core、Extension Core 和 Model Core 主链路已经完成;第二阶段 Agent Trace 持久化、可恢复 SSE、stdio MCP Bridge、隔离 Plugin Host 以及 Plugin Command/Settings 已落地,后端当前回归基线为 136 项测试通过。
|
||||
|
||||
## 当前实现
|
||||
|
||||
@@ -342,11 +342,11 @@ Skill Manifest
|
||||
|
||||
前端智能体页面已经完成中文联调:运行状态、Agent Event、内置 Tool、Permission 和常用事件详情字段均通过集中标签映射展示中文;`notes.search` 等技术 ID 继续保留,便于与后端 Trace、日志和接口契约对应。
|
||||
|
||||
- 已实现 Mock、OpenAI-Compatible Chat Completions 与 Ollama Adapter;OpenAI Responses 和 Anthropic Messages 尚未实现。
|
||||
- 已实现 Mock、OpenAI-Compatible Chat Completions、Ollama、OpenAI Responses 和 Anthropic Messages Adapter;阶段 E 同时完成国内预设、持久化配置和能力模型路由,详见 [模型提供商与模型发现开发说明](模型提供商与模型发现开发说明.md)。
|
||||
- Provider 配置暂存内存,后续通过 Repository 接入 SQLite;PATCH 已支持用显式 `null` 清空 base URL、默认模型和凭据引用。
|
||||
- Run/Trace 已通过 Repository 接入 SQLite;后续增加按保留策略归档和 Benchmark 引用保护。
|
||||
- Permission 已有核心等待/恢复机制,前端确认 UI 已完成联调和中文展示。
|
||||
- Task 已持久化到 SQLite;Attachment Tool 读取 Host 管理目录中的 UTF-8 文件。
|
||||
- `audio.transcribe` 当前消费 Host 预生成的 transcript;faster-whisper 与说话人分离仍按技术基线在第二阶段接入。
|
||||
- `audio.transcribe` 当前消费 Host 预生成的 transcript;faster-whisper 与说话人分离仍待第二阶段后续接入。
|
||||
- Extension 安装记录暂存内存;后续接入持久化 Registry 与版本升级流程。
|
||||
- 当前 Plugin Host 支持内置声明式 handler、本地 stdio MCP Server 以及 Plugin Command/Settings;Streamable HTTP、OS 级沙箱与 UI Contribution 留在后续阶段。
|
||||
- 当前 Plugin Host 支持内置声明式 handler、本地 stdio MCP Server 以及 Plugin Command/Settings;独立 MCP Server Registry 另行支持 stdio、Streamable HTTP 与旧 SSE 兼容。OS 级沙箱与 UI Contribution 留在后续阶段。
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
# Benchmark 开发说明
|
||||
|
||||
> 所属模块:Knowledge / Retrieval Core(后端,负责人 yxx)。RAG Benchmark 已交付;Agent Benchmark 暂缓,待 Agent Runtime 完成后在同一契约下补齐。
|
||||
|
||||
## 定位
|
||||
|
||||
Benchmark Service 用受控 Dataset 对检索引擎做可复现评测:创建即返回 queued、后台 asyncio.Task 执行、SSE 实时推送进度、结束后产出结构化报告。CLI、测试与前端报告页复用同一 Service,不各自实现指标。
|
||||
|
||||
## 接口
|
||||
|
||||
| 方法 | 路径 | 用途 |
|
||||
| --- | --- | --- |
|
||||
| GET | `/api/benchmarks/datasets?kind=rag` | 枚举受控目录下的 Dataset 元信息 |
|
||||
| POST | `/api/benchmarks/rag/runs` | 创建 RAG Benchmark(202) |
|
||||
| GET | `/api/benchmarks/runs?kind=&status=&limit=&offset=` | 分页获取运行记录 |
|
||||
| GET | `/api/benchmarks/runs/{run_id}` | 状态与指标摘要 |
|
||||
| GET | `/api/benchmarks/runs/{run_id}/events` | SSE 进度与 Case 结果 |
|
||||
| POST | `/api/benchmarks/runs/{run_id}/cancel` | 取消运行 |
|
||||
| GET | `/api/benchmarks/runs/{run_id}/report` | 结构化完整报告 |
|
||||
|
||||
Agent Benchmark 的 `/api/benchmarks/agent/runs` 未暴露(暂缓),不在 OpenAPI 注册占位接口。
|
||||
|
||||
## Dataset
|
||||
|
||||
Dataset 来自 `settings.benchmark_datasets_path`(默认 `backend/data/benchmarks`),API 不接受调用方提交任意路径。按文件名 stem 精确匹配 `{dataset_id}.json`,与请求无关文件的损坏(JSON 语法错误、UTF-8 解码错误、顶层非对象)不会阻断加载;只有目标文件本身损坏才返回 `BENCHMARK_DATASET_INVALID`。
|
||||
|
||||
RAG Case 结构:`case_id`、`query`、`expected_note_ids`、`expected_block_ids`、`citation_required`、`tags`。`citation_required=true` 时必须声明 `expected_block_ids`,否则无法计算 Citation Hit Rate。
|
||||
|
||||
## 运行生命周期
|
||||
|
||||
`queued → running → completed | failed | cancelled`。
|
||||
|
||||
- 创建时校验索引兼容性:索引非空、Embedding model/dim 与当前引擎一致、vector/hybrid 时向量索引非空;不满足返回 `BENCHMARK_INDEX_INCOMPATIBLE`(409),避免把环境/索引错误误判为检索质量差。
|
||||
- 内存注册表上限 `MAX_RUNS=100`,超限只淘汰终态 run;满容量且全为活动 run 时返回 `BENCHMARK_CAPACITY_EXCEEDED`(429)。
|
||||
- 失败/取消只向公开响应暴露项目错误码与安全消息,详细异常进入日志,不通过 HTTP/SSE 返回。
|
||||
|
||||
## 指标
|
||||
|
||||
RAG 按 (mode, case, repeat) 逐样本计算,再按 mode 聚合:
|
||||
|
||||
- 质量:`hit_at_1`、`hit_at_5`、`recall_at_k`、`mrr`、`citation_hit_rate`;
|
||||
- 延迟:`p50_latency_ms`、`p95_latency_ms`(仅统计成功样本);
|
||||
- 样本构成:`total_cases`、`successful_cases`、`failed_cases`、`failure_rate`。
|
||||
|
||||
失败样本按零分计入质量指标分母,报告据此可知实际分母,避免把执行失败误判为检索质量差。
|
||||
|
||||
## 事件与 SSE
|
||||
|
||||
事件流:`RunStarted → CaseCompleted* → RunCompleted | RunFailed | RunCancelled`。
|
||||
|
||||
`GET /api/benchmarks/runs/{run_id}/events` 支持 `Last-Event-ID` 与 `?after_sequence=` 游标恢复(复用 Agent SSE 的解析逻辑),`RunCompleted` / `RunFailed` / `RunCancelled` 为终止事件,收到后断流。
|
||||
|
||||
## 错误码
|
||||
|
||||
```text
|
||||
BENCHMARK_DATASET_NOT_FOUND
|
||||
BENCHMARK_DATASET_INVALID
|
||||
BENCHMARK_INDEX_INCOMPATIBLE
|
||||
BENCHMARK_CAPACITY_EXCEEDED
|
||||
BENCHMARK_RUN_NOT_FOUND
|
||||
BENCHMARK_RUN_FAILED
|
||||
BENCHMARK_CASE_EVALUATION_FAILED
|
||||
```
|
||||
|
||||
## 配置快照
|
||||
|
||||
报告与运行记录保存 `config_snapshot`:dataset hash/version、modes、retrieval 参数、Reranker、索引元数据、App 版本与环境、Python 版本。`local_embedding` 记录本地基线 model/version/dim;`embedding.policy = per_case` 表示实际来源以逐样本结果为准,不能把本地基线当作本次使用的模型。
|
||||
|
||||
每个 `RAGCaseResult.embedding`(同时出现在报告 cases 和 CaseCompleted SSE 中)记录 `source`(api/local/not_used/unavailable)、实际 `model_id` 空间标识、`dimensions`、本地 `version`、`fallback_reason`。远程路由还记录请求时的 `route_version` 和 `requested_route`(provider_id/model/endpoint/dimensions,不含凭据)、成功生成查询向量后的 `attempted_space`。FTS 标记 not_used;调用失败而未完成向量检索时标记 unavailable。API 不可用或远程索引缺失时,实际模型仍记录最终使用的本地基线。配置允许在样本间改变,逐样本记录对应实际调用;汇总指标可能包含多种空间,比较实验时需检查 cases。记录使用任务局部上下文隔离,并发评测不会相互覆盖。
|
||||
|
||||
## 测试
|
||||
|
||||
```powershell
|
||||
cd backend
|
||||
uv run pytest -q
|
||||
```
|
||||
|
||||
`tests/test_benchmark.py` 覆盖数据集注册与校验、指标纯函数、端到端运行、取消、索引兼容、容量与失败样本聚合;`tests/test_retrieval.py` 覆盖 FTS 阈值与分页 total 一致性。
|
||||
@@ -3,7 +3,7 @@
|
||||
> 本文档用于团队开发和模块联调,记录 Knowledge Core / Retrieval Core 已经落地的
|
||||
> 模块边界、数据模型、接口与使用方式,对应分工表中的杨星萱。
|
||||
|
||||
> 更新日期:2026-09-02。第一阶段 Knowledge/Retrieval 主链路已经完成,并已接入 Agent Tool Registry;完整后端回归基线为 126 项测试通过。
|
||||
> 更新日期:2026-09-02。第一阶段 Knowledge/Retrieval 主链路已经完成,并已接入 Agent Tool Registry;完整后端回归基线为 136 项测试通过。
|
||||
|
||||
## 当前实现
|
||||
|
||||
@@ -198,7 +198,7 @@ cd backend
|
||||
uv run pytest -q
|
||||
```
|
||||
|
||||
当前后端完整测试共 71 个用例通过(单元 + 端到端)。测试通过 `tests/conftest.py` 的 autouse fixture 把
|
||||
当前后端完整测试共 218 个用例通过(单元 + 端到端)。测试通过 `tests/conftest.py` 的 autouse fixture 把
|
||||
数据目录/DB/Vault 重定向到临时目录,不读写真实 `backend/data`,任何本机状态下结果确定。
|
||||
|
||||
## 配置
|
||||
@@ -232,4 +232,6 @@ rag.search
|
||||
- Embedding / Reranker 为轻量实现,后续替换为真实模型(接口不变)。
|
||||
- 小语料下 hybrid 检索召回偏宽(向量 Top-K 覆盖全部 block),可加相关性阈值收紧。
|
||||
- 重建为同步 + 全量,后续接入增量索引与异步任务队列。
|
||||
- 检索 Benchmark 待建立。
|
||||
- RAG Benchmark 已建立:`POST /api/benchmarks/rag/runs` 创建即返回 queued、后台 Task 执行,
|
||||
通过 SSE 实时推送进度,报告含逐 Case 结果与 `total_cases` / `successful_cases` / `failed_cases` / `failure_rate`。
|
||||
- Agent Benchmark 暂缓,待 Agent Runtime 完成后交付。
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# Plugin Command 与 Settings 开发说明
|
||||
|
||||
> 更新日期:2026-09-02。本文记录第二阶段阶段 D 已实现的 Plugin Command Contribution、Plugin Settings Contribution、Secret 边界和前端 Service Contract。当前回归基线为后端 126 项测试、前端 29 项测试,TypeScript 类型检查和生产构建通过。
|
||||
> 更新日期:2026-09-03。本文记录第二阶段阶段 D 已实现的 Plugin Command Contribution、Plugin Settings Contribution、Secret 边界及前端闭环。当前回归基线为后端 136 项测试、前端 32 项测试,TypeScript 类型检查和生产构建通过。
|
||||
|
||||
## 1. 阶段目标
|
||||
|
||||
@@ -8,9 +8,9 @@
|
||||
|
||||
- Command:插件声明命令,宿主负责注册、展示、校验、执行和返回白名单 effect;
|
||||
- Settings:插件声明设置 Schema,宿主负责动态表单 Contract、非敏感值持久化和 Secret 加密引用;
|
||||
- Frontend Contract:提供稳定的 TypeScript DTO 与 Service,供后续命令面板、右键菜单和插件设置页直接联调。
|
||||
- Frontend:Plugin 详情页提供 Host 状态、动态设置、Secret 管理与命令执行,全局命令面板加载 command_palette Contribution。
|
||||
|
||||
本阶段不实现前端页面,也不把第三方代码导入 FastAPI 进程。操作系统级安全沙箱仍按规划在第三阶段桌面基础集成完成后、Tauri/Rust 沙箱正式构建前处理。
|
||||
第三方代码不会导入 FastAPI 进程。操作系统级安全沙箱仍按规划在第三阶段桌面基础集成完成后、Tauri/Rust 沙箱正式构建前处理。
|
||||
|
||||
## 2. 包内声明
|
||||
|
||||
@@ -86,7 +86,9 @@ PUT /api/plugins/{plugin_id}/settings/{key}/secret
|
||||
DELETE /api/plugins/{plugin_id}/settings/{key}/secret
|
||||
```
|
||||
|
||||
前端 `pluginService` 已提供对应方法及 Wire DTO,但阶段 D 不创建命令面板或动态设置表单页面。调用方必须使用服务层,不自行拼接路径;Secret 不得写入 Pinia、LocalStorage 或调试日志。
|
||||
前端 `pluginService` 提供对应方法及 Wire DTO,Plugin 详情页据此展示 MCP Host 状态和重启入口、动态生成五类设置字段、独立写入或删除 Secret,并执行带参数的 Plugin Command。全局命令面板打开时获取 `command_palette` 命令;需要必填参数的命令会引导用户进入详情页填写。调用方必须使用服务层,不自行拼接路径。
|
||||
|
||||
Secret 明文仅存在于当前密码输入框绑定的组件内存,提交后立即清空;不得写入 Pinia、LocalStorage、普通 Settings 请求或调试日志。前端不会读取 Secret 明文,只展示后端返回的 `configured` 状态。
|
||||
|
||||
## 6. 主要错误边界
|
||||
|
||||
@@ -110,6 +112,6 @@ pnpm type-check
|
||||
pnpm build
|
||||
```
|
||||
|
||||
阶段 D 测试覆盖注册/注销生命周期、位置过滤、参数与 Context 校验、上下文裁剪、设置影响命令执行、声明式 Secret Resolver 与越权拒绝、真实 MCP Command Target 与 Agent Tool 隔离、必填 Secret 传递、外部 Schema 引用拒绝、定长 Secret Reference、篡改引用的跨命名空间阻断、Secret 删除与卸载失败回滚、Provider/通用凭据命名空间隔离、五类设置字段、Schema 版本冲突、Secret 密文与清理、损坏存储、空 Command 列表等无效贡献文件、OpenAPI 路径和前端 Service 请求格式。
|
||||
阶段 D 测试覆盖注册/注销生命周期、位置过滤、参数与 Context 校验、上下文裁剪、设置影响命令执行、声明式 Secret Resolver 与越权拒绝、真实 MCP Command Target 与 Agent Tool 隔离、必填 Secret 传递、外部 Schema 引用拒绝、定长 Secret Reference、篡改引用的跨命名空间阻断、Secret 删除与卸载失败回滚、Provider/通用凭据命名空间隔离、五类设置字段、Schema 版本冲突、Secret 密文与清理、损坏存储、空 Command 列表等无效贡献文件、OpenAPI 路径、前端 Service 请求格式、Host 状态展示和动态 Secret 表单。
|
||||
|
||||
生产构建仍会报告现有大 Chunk 警告,不影响构建成功;该问题属于前端按路由和 Markdown 依赖拆包的后续性能任务。
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# 前端写作体验优化开发说明
|
||||
|
||||
> 更新日期:2026-08-30。本文所述优化均已进入当前分支;前端完整回归基线为 14 项测试通过,TypeScript 检查和 Vite 生产构建通过。
|
||||
> 更新日期:2026-09-02。本文所述优化均已进入 `main`;当前前端完整回归基线为 29 项测试通过,TypeScript 检查和 Vite 生产构建通过。
|
||||
|
||||
## 1. 本次目标
|
||||
|
||||
@@ -103,7 +103,7 @@ pnpm build
|
||||
pnpm test
|
||||
```
|
||||
|
||||
验证结果:TypeScript 类型检查与 Vite 生产构建均通过。当前前端完整回归测试共 14 项;其中写作与文件切换相关回归覆盖:
|
||||
验证结果:TypeScript 类型检查与 Vite 生产构建均通过。当前前端完整回归测试共 29 项;其中写作与文件切换相关回归覆盖:
|
||||
|
||||
- 顶部工具栏对选区应用加粗;
|
||||
- 浮动工具栏对选区应用斜体;
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# 前端壳子与接口层开发说明
|
||||
|
||||
> 更新日期:2026-08-30
|
||||
> 更新日期:2026-09-02
|
||||
> 适用范围:Vue 3 + TypeScript 页面、Workspace、公共 Service、FastAPI 接口适配和 SSE。
|
||||
> 文档用途:帮助团队理解当前前端可用能力、模块边界、启动方式和后续页面开发入口。
|
||||
|
||||
@@ -187,12 +187,12 @@ pnpm build
|
||||
```text
|
||||
pnpm build passed
|
||||
pnpm test 29 passed
|
||||
uv run pytest 126 passed
|
||||
uv run pytest 136 passed
|
||||
preview smoke HTTP 200
|
||||
git diff --check passed
|
||||
```
|
||||
|
||||
当前前端使用 Vitest 执行 Store、Workspace API Adapter、SSE 恢复游标、Plugin Command/Settings Service、文件树、编辑器组件、智能体标签、轻量动效约束、Markdown 对比度 Token、scoped CSS 选择器约束和 Shiki GitHub 双主题测试;`pnpm build` 同时执行 `vue-tsc -b` 与 Vite 生产构建。后端测试出现过 `.pytest_cache` 无法写入的 Windows 权限警告,不影响 126 项测试结果,也不涉及产品代码。
|
||||
当前前端使用 Vitest 执行 Store、Workspace API Adapter、SSE 恢复游标、Plugin Command/Settings Service、文件树、编辑器组件、智能体标签、轻量动效约束、Markdown 对比度 Token、scoped CSS 选择器约束和 Shiki GitHub 双主题测试;`pnpm build` 同时执行 `vue-tsc -b` 与 Vite 生产构建。后端测试出现过 `.pytest_cache` 无法写入的 Windows 权限警告,不影响 136 项测试结果,也不涉及产品代码。
|
||||
|
||||
Vite 当前会提示 Chat 与 Workspace 的部分异步 Chunk 超过 500 kB,这是 Milkdown、CodeMirror、KaTeX 和 Shiki 等编辑/渲染依赖带来的性能优化项,不影响构建成功或功能正确性;进入桌面打包前应通过手动分包或更细粒度动态加载继续优化。
|
||||
|
||||
|
||||
@@ -1,107 +1,107 @@
|
||||
# 模型提供商与模型发现开发说明
|
||||
# 模型提供商、协议适配与模型路由开发说明
|
||||
|
||||
> 更新日期:2026-08-30。OpenAI、DeepSeek、Ollama 预设、模型自动发现、默认模型选择和开发阶段加密凭据存储均已实现并接入设置页。
|
||||
> 更新日期:2026-09-04。阶段 E 实现记录。本地小模型的实际安装与多模态队列属于阶段 F;本阶段保留并测试可注入的本地后端接口。
|
||||
|
||||
## 1. 本次目标
|
||||
## 1. 设置与凭据
|
||||
|
||||
本次完善设置页的模型提供商配置,不改变 Agent、Chat 和 Skill 对统一 Model Core 接口的依赖:
|
||||
设置 → 模型提供商 → 新增 Provider 提供可搜索的 logo 预设网格,包含 DeepSeek、Kimi、阿里云百炼、智谱 GLM、火山方舟、硅基流动、百度千帆、腾讯混元、MiniMax、阶跃星辰,以及 OpenAI Chat / Responses、Anthropic 和 Ollama。图标打包到前端,使用时不请求第三方图片服务;来源和许可见前端 assets/providers 目录。
|
||||
|
||||
- 提供 OpenAI、DeepSeek 和 Ollama 配置预设;
|
||||
- 保存 Provider 后自动获取该账号或服务当前可用的模型列表;
|
||||
- 支持手动刷新模型列表和选择默认模型;
|
||||
- 保留自定义 OpenAI-Compatible 服务入口;
|
||||
- 不在 Vue、FastAPI 配置或仓库文件中保存、回显 API Key 明文。
|
||||
预设返回 `preset_id`、`logo_id`、`name`、`provider_type`、`base_url`、`requires_credential`、`description` 和 `capabilities`。能力标签表示预设接入范围,不保证该账号的每个模型支持全部能力。厂商专用媒体协议、Coding Plan 和海外地域需要使用对应地址,不能仅凭厂商名称推断协议兼容。
|
||||
|
||||
## 2. 接口与实现
|
||||
预设和自定义服务都可以直接输入 API Key。每个新配置分配独立 Credential ID,避免同厂商多账号相互覆盖。明文只留在密码输入框和专用请求中,提交、失败、切换预设及关闭时清空;密钥不进入 Pinia、localStorage、Provider 配置响应或模型路由。
|
||||
|
||||
### 2.1 Provider 预设
|
||||
凭据继续使用独立的 `PUT /api/credentials/{credential_id}` 和 Fernet 开发存储。`plugin.*`、`mcp.*` 是保留命名空间。桌面端阶段仍需要把主密钥管理迁移到 Stronghold。保存密钥与保存 Provider 是两个请求,Provider 保存失败时可能留下未引用的加密凭据,可通过凭据删除接口清理。
|
||||
|
||||
新增接口:
|
||||
Provider 配置和 Credential ID 写入 SQLite `provider_configs`,重启后恢复。Mock 为内置 Provider,不能编辑或删除。PATCH 已支持变更 `provider_type` 并重新创建 Adapter;Base URL 限制为不带用户信息、查询或 fragment 的 HTTP(S) 地址。
|
||||
|
||||
```http
|
||||
GET /api/providers/presets
|
||||
## 2. 协议适配
|
||||
|
||||
支持的协议是 OpenAI Chat Completions、OpenAI-Compatible、OpenAI Responses、Anthropic Messages 和 Ollama。Agent、Chat、Skill 仍只依赖内部 `ModelRequest` / `ModelEvent` / `ProviderTurn`,不直接解释厂商协议。
|
||||
|
||||
Adapter 负责消息及 Tool 历史转换、增量文本、可用的 reasoning delta、工具参数片段、usage、终止与统一错误。外部错误正文不原样返回;HTTP 鉴权、限流、超时、无效数据、流中断分别映射为内部错误。取消继续传播并关闭上游连接,不触发第二次本地推理。
|
||||
|
||||
`GET /api/providers/{provider_id}/models` 用于发现模型。模型列表不等于每个模型的能力承诺;部分厂商或代理不提供 `/models` 时,允许直接手动输入模型 ID。连接测试验证模型发现接口,不代表每一种媒体模型已完成真实推理验收。
|
||||
|
||||
## 3. 三类模型路由
|
||||
|
||||
接口:
|
||||
|
||||
| 方法 | 路径 | 用途 |
|
||||
| --- | --- | --- |
|
||||
| GET | `/api/model-routing` | 读取配置和本地后端状态 |
|
||||
| PUT | `/api/model-routing` | 带版本更新三类模型绑定 |
|
||||
| POST | `/api/models/embeddings` | 文本向量,返回来源和回退原因 |
|
||||
| POST | `/api/media/transcriptions` | 附件转写作业 |
|
||||
| GET | `/api/media/transcriptions/{job_id}` | 获取转写作业 |
|
||||
| POST | `/api/media/speaker-matches` | 两个音频附件的声纹相似度 |
|
||||
|
||||
设置 → 索引与模型分别选择 Embedding、音频转文本和声纹匹配。三种绑定互相独立,可使用不同提供商、模型、密钥和 API 路径。
|
||||
|
||||
GET / PUT 响应:
|
||||
|
||||
```json
|
||||
{
|
||||
"config": {
|
||||
"version": 1,
|
||||
"embedding": {
|
||||
"provider_id": "provider_example",
|
||||
"model": "your-embedding-model",
|
||||
"endpoint": "/embeddings",
|
||||
"dimensions": null
|
||||
},
|
||||
"transcription": null,
|
||||
"speaker_matching": null
|
||||
},
|
||||
"local_backends": [
|
||||
{"capability": "embedding", "status": "placeholder", "message": "当前为 hash-v1 占位向量"},
|
||||
{"capability": "transcription", "status": "not_installed", "message": "阶段 F 接入"},
|
||||
{"capability": "speaker_matching", "status": "not_installed", "message": "阶段 F 接入"}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
预设由后端 `ProviderFactory` 提供,前端只消费名称、协议类型、Base URL 和是否需要凭据等配置元数据,不直接实现厂商协议。
|
||||
PUT body 只提交 `config` 的内容。`version` 为读取时的版本,成功递增;并发更新返回 `MODEL_ROUTING_VERSION_CONFLICT`。绑定为空表示使用本地后端。删除仍被路由引用的 Provider 返回 `PROVIDER_IN_USE`,须先解除绑定。
|
||||
|
||||
当前预设:
|
||||
本阶段三类远程路由使用 `openai_chat` / `openai_compatible` 的 Bearer HTTP 配置,endpoint 只能是该提供商下的路径。Responses、Anthropic 和 Ollama 原生协议不冒充上述媒体协议;Ollama 用户需要另建兼容 HTTP 配置才能用于当前远程 Embedding 接口。
|
||||
|
||||
| 提供商 | Provider Type | Base URL | 默认 Credential ID |
|
||||
| --- | --- | --- | --- |
|
||||
| OpenAI | `openai_chat` | `https://api.openai.com/v1` | `openai` |
|
||||
| DeepSeek | `openai_compatible` | `https://api.deepseek.com` | `deepseek` |
|
||||
| Ollama | `ollama` | `http://127.0.0.1:11434` | 无 |
|
||||
调用规则:无绑定 → 本地接口;有绑定 → API → 校验结果 → 失败或无效时调用本地接口。Provider 停用、密钥缺失、鉴权失败、限流、网络超时及无效结果均可回退;用户取消不会回退。附件不存在、大小非法等输入错误直接返回,不把用户输入错误当成模型故障。
|
||||
|
||||
OpenAI 和 DeepSeek 都通过项目已有的 `OpenAICompatibleProvider` 访问。模型发现分别请求 Base URL 下的 `/models`,不引入厂商 SDK。
|
||||
## 4. Embedding 与索引一致性
|
||||
|
||||
### 2.2 自动获取模型
|
||||
请求使用 `model`、`input`、`encoding_format: float`;只有明确配置维度时才发送 `dimensions`。按最多 32 条分批请求,全部批次有效才使用 API 结果。校验返回数量、连续唯一 index、维度一致性、有限数值、非零范数,并 L2 归一化。维度可为 1–16384,不截断、补零或混用不同模型的向量。
|
||||
|
||||
模型列表继续使用既有接口:
|
||||
返回 `vectors`、`source`、`model_id`、`dimensions`、`fallback_reason`。远程空间 ID 由完整 API URL、模型和实际维度生成;即使维度相同,不同模型的空间也不同。
|
||||
|
||||
```http
|
||||
GET /api/providers/{provider_id}/models
|
||||
```
|
||||
笔记索引始终保留现有 hash/sqlite-vec 本地基线,远程向量写入独立 `routed_block_vectors` 表。远程查询只搜索对应空间,并要求覆盖全部当前 Block。API 失败、索引缺失、不完整或损坏时使用完整本地索引。切换模型、URL、维度后应在设置中重建全部索引。旧空间与当前文本不会混合打分,删除笔记或重建索引会通过外键清理远程向量。
|
||||
|
||||
设置页在以下时机调用该接口:
|
||||
当前远程侧索引采用 SQLite JSON 向量和精确余弦扫描,复杂度 O(Block 数量 × 维度),适用于当前小型 Vault;后续大规模索引需替换为按空间隔离的 ANN。网络等待发生在数据库写事务之前,当前仍会增加保存或重建延迟,异步索引队列尚未接入。全量重建先在内存中准备全部向量,再使用一个 SQLite 事务更新元数据、FTS、本地与远程向量及任务关联;取消或失败只回滚索引事务,不再覆盖整库文件。准备阶段保留旧索引可查询,代价是内存同时容纳本次重建的向量。
|
||||
|
||||
- Provider 列表加载完成后,为所有已启用 Provider 自动刷新;
|
||||
- 新增或编辑 Provider 保存成功后自动刷新;
|
||||
- 用户点击“刷新模型”时手动刷新;
|
||||
- 打开已有 Provider 的编辑窗口时刷新可选模型。
|
||||
OpenAI Compatible 流中,工具名称可能分片返回。适配器在本轮输出结束后发送完整工具名及已缓冲参数,避免把名称片段当作工具 ID;文本与推理内容仍逐片发送。
|
||||
|
||||
前端按模型名称排序并按 `model_id` 去重。获取结果保存在 `providerStore.modelsByProvider`,加载状态和错误按 Provider 隔离,单个外部服务失败不会阻止其他服务展示。
|
||||
无 API 时使用的 `HashEmbeddingProvider` 是确定性特征哈希占位实现,**不是已集成的小型语义模型**。真实本地 Embedding 可实现既有 `EmbeddingProvider` 接口注入。
|
||||
|
||||
获取成功后,Provider 卡片展示模型数量和默认模型下拉框。更换默认模型会调用 Provider PATCH 接口写回配置;编辑窗口仍允许手动输入模型 ID,以兼容未出现在列表中的代理模型或部署别名。
|
||||
## 5. 音频与声纹边界
|
||||
|
||||
### 2.3 错误处理
|
||||
转写默认请求 `/audio/transcriptions`,multipart 字段 `model`、可选 `language` 和 `file`,响应必须包含非空字符串 `text`。已有纯文本附件和 Host 旁路 `.txt` 导入保留,来源标记 `sidecar`,不伪称 ASR。转写作业新增 `source`、`fallback_reason`;回退失败的作业记录 `LOCAL_MODEL_NOT_INSTALLED` 等明确错误。作业目前同步执行、限量保存在内存中,不是持久化异步队列。
|
||||
|
||||
Provider Adapter 的错误在 FastAPI 路由转换为统一 API Error:
|
||||
声纹匹配使用**本项目自定义 HTTP 契约**,默认 `/audio/speaker-matches`,multipart 字段 `model`、`file`、`reference_file`;响应为 `{"score": 0.85}`,score 必须为有限的 0–1 数值。公共入口只接受 `attachment_id` 和 `reference_attachment_id`,不接收任意文件路径。此接口用于一对一声纹比对,不等同于 pyannote 说话人分离,也不声称任意国内厂商原生支持该路径。
|
||||
|
||||
| Provider Error | HTTP 状态 |
|
||||
| --- | --- |
|
||||
| `PROVIDER_AUTH_FAILED` | 401 |
|
||||
| `MODEL_NOT_FOUND` | 404 |
|
||||
| `PROVIDER_RATE_LIMITED` | 429 |
|
||||
| `PROVIDER_TIMEOUT` | 504 |
|
||||
| 其他 Provider 可用性错误 | 502 |
|
||||
媒体文件限制 1 字节至 25 MiB,API 响应限制 16 MiB,单次请求超时 30 秒。文件从后端受控附件目录读取,使用结束或取消时关闭句柄。
|
||||
|
||||
前端在对应 Provider 卡片内展示失败原因,并允许用户修正 Credential ID、Base URL 后重新获取。
|
||||
`LocalSpeechBackend` 提供 `transcribe` 和 `match` 接口。阶段 E 默认 `PendingSpeechBackend` 明确报告未安装;阶段 F 接入 faster-whisper、pyannote.audio 及模型资源后替换。当前 `diarization=true` 明确返回失败作业 `DIARIZATION_NOT_IMPLEMENTED`,不会静默忽略。视频解码、TTS、视频生成及厂商专用异步媒体协议不在本次交付内。
|
||||
|
||||
## 3. 凭据边界
|
||||
## 6. 官方协议依据与验证
|
||||
|
||||
设置页选择 OpenAI 或 DeepSeek 预设后展示密码类型的 API Key 输入框,不再要求用户理解 Credential ID。输入值只存在于表单的临时 `ref`,不会写入 Pinia 或 localStorage;请求完成、取消表单或失败后都会清空。
|
||||
国内通用地址核对依据:[阿里云百炼兼容接口](https://help.aliyun.com/zh/model-studio/compatibility-of-openai-with-dashscope)、[百度千帆兼容接口](https://cloud.baidu.com/doc/qianfan/s/Hmh4suq26)、[腾讯混元兼容接口](https://cloud.tencent.com/document/product/1729/111007)、[MiniMax 文本接口](https://platform.minimaxi.com/docs/guides/text-generation)、[阶跃星辰通用与套餐地址区别](https://platform.stepfun.com/docs/zh/step-plan/overview)、[火山方舟 API](https://www.volcengine.com/docs/82379/1795150)、[智谱开放接口](https://docs.bigmodel.cn/api-reference/文件-api/文件列表)。模型 ID 以账号实际开通列表为准,不写死“最新模型”。
|
||||
|
||||
API Key 通过独立接口写入:
|
||||
流式事件依据:[OpenAI Responses streaming](https://platform.openai.com/docs/api-reference/responses-streaming)、[Anthropic streaming](https://platform.claude.com/docs/en/build-with-claude/streaming)。音频请求依据:[SiliconFlow transcription](https://docs.siliconflow.com/en/api-reference/audio/create-audio-transcriptions)。
|
||||
|
||||
```http
|
||||
GET /api/credentials/{credential_id}
|
||||
PUT /api/credentials/{credential_id}
|
||||
DELETE /api/credentials/{credential_id}
|
||||
```
|
||||
自动化验证使用虚构凭据、本地附件、httpx.MockTransport 和可注入本地模型,覆盖流式 Tool/Usage/取消、错误映射、回退、索引空间隔离、版本冲突、重启恢复和界面凭据行为。没有使用真实 API Key 或向厂商发送推理请求。审阅修复并同步主分支后验证:后端全量 447 项、前端 76 项测试通过,Vue/TypeScript 类型检查和生产构建通过,浅色/深色预设页面与路由保存经过浏览器检查,git diff --check 通过。后端仅保留既有 Starlette 测试客户端弃用提示,前端保留既有大 bundle 提示。
|
||||
|
||||
PUT 请求使用 Pydantic `SecretStr` 接收密钥,响应仅包含 Credential ID 和 `configured` 状态。后端使用 Fernet 认证加密,将密文保存到 `data/credentials/credentials.json`,主密钥保存到 `data/credentials/master.key`;目录和文件尽可能设置为仅当前用户可访问并整体排除版本控制。写入采用临时文件替换,避免进程中断留下半写文件。Provider 发起请求时按 Credential ID 解密,解密失败转换为统一 Provider Error,任何读取接口均不返回明文。
|
||||
|
||||
本地开发存储的主密钥与密文仍位于同一用户数据目录,因此它解决的是仓库泄漏、普通配置误提交和静态明文暴露,不等同于操作系统安全硬件或 Stronghold。Tauri 集成后应以 Stronghold 实现替换 `EncryptedCredentialStore`。无界面环境仍兼容 `OPENAI_API_KEY`、`DEEPSEEK_API_KEY` 和 Host 注入的 `AINOTE_CREDENTIAL_<ID>`;设置页保存的本地密钥优先,环境变量仅作为回退。
|
||||
|
||||
自动化测试仅使用虚构测试值,验证磁盘文件不包含明文、加解密往返、API 响应不泄密,以及 Provider 能用解密后的值构造 Authorization Header。本次没有使用真实 OpenAI 或 DeepSeek Key,也没有向厂商发起真实请求。
|
||||
|
||||
## 4. 验证
|
||||
|
||||
后端:
|
||||
|
||||
```bash
|
||||
```powershell
|
||||
cd backend
|
||||
uv run pytest -q -p no:cacheprovider
|
||||
```
|
||||
|
||||
前端:
|
||||
|
||||
```bash
|
||||
cd frontend
|
||||
cd ../frontend
|
||||
pnpm test
|
||||
pnpm build
|
||||
```
|
||||
|
||||
自动化验证覆盖 Provider 预设、OpenAI-Compatible `/models` 请求与鉴权头、模型映射、前端自动刷新、排序去重及按 Provider 隔离错误。生产构建同时执行 Vue 和 TypeScript 类型检查。
|
||||
|
||||
当前完整回归基线:后端 126 项测试、前端 29 项测试通过,前端类型检查和生产构建通过。Provider 配置目前仍保存在内存 Registry,AI Core 重启后需要重新创建;凭据密文会保留。`plugin.*` 为 Plugin Secret 保留命名空间,Provider 配置、临时测试凭据和通用凭据 API 均拒绝该前缀。OpenAI Responses 与 Anthropic Messages Adapter 尚未实现,设置页正式预设不会使用这两种协议。
|
||||
|
||||
@@ -0,0 +1,141 @@
|
||||
# 独立 MCP Server 配置中心开发说明
|
||||
|
||||
> 更新日期:2026-09-03。本文记录第二阶段 C.1 的完整实现;独立 MCP Server Registry 与 Plugin 自带 MCP Host 是两个并列入口。
|
||||
|
||||
## 1. 已实现范围
|
||||
|
||||
- 独立 Server 的创建、读取、版本化编辑、删除和 Tool 摘要查询;
|
||||
- `stdio`、Streamable HTTP 和旧版 HTTP+SSE 三种 Transport;
|
||||
- stdio 可执行文件、参数、普通/加密环境变量,以及 HTTP URL、普通/加密 Header;
|
||||
- 配置摘要确认、连接测试、启停、异常状态与最近一次测试结果;
|
||||
- MCP initialize、`tools/list`、`tools/call`、取消与动态 Tool 注册,名称为 `mcp.{server_id}.{tool}`;
|
||||
- Streamable HTTP Session、协议版本 Header、JSON/SSE POST 响应、可选 GET 事件流和 `Last-Event-ID` 重连;
|
||||
- 旧 HTTP+SSE 的 endpoint 事件与消息 POST,并强制消息地址和配置地址同源;
|
||||
- 前端表单/JSON 双模式、三种模板、高风险变更确认及请求期 Secret 输入。
|
||||
|
||||
Streamable HTTP 按 MCP 当前规范实现;SSE 仅用于兼容旧 Server,不应作为新部署首选。
|
||||
|
||||
## 2. 配置、版本与 Secret
|
||||
|
||||
普通配置原子写入 `APP_DATA_DIR/mcp/servers.json`。更新请求必须携带读取到的 `version`;版本过期返回 `409 MCP_SERVER_VERSION_CONFLICT`,避免多个页面互相覆盖。改变 Transport、命令、URL、Header、环境变量或权限后,旧授权和测试结果立即失效。
|
||||
|
||||
`backend/data/mcp/` 是本机运行数据,包含连接配置、授权状态和第三方进程工作目录,不属于团队共享配置。`.gitignore` 忽略整个目录以及 `server.json`、`servers.json` 文件名;不得强制添加到 Git。提交前检查暂存文件清单,不要将本地密钥、连接配置或运行数据推送到远程。
|
||||
|
||||
Secret 使用带类型和键名哈希的 `mcp.*` 内部 ID 写入 Fernet 凭据存储。查询 API 只返回环境变量或 Header 是否配置,不返回明文。删除 Server 或移除 Secret 键会同步清理密文。前端密码框提交后立即清空;用户主动粘贴到 JSON 的密钥仅在当前编辑会话中暂存,解析后从 JSON 中移除,不写入 localStorage、普通配置或日志。
|
||||
|
||||
环境变量的凭据 ID 使用区分大小写的 v2 名称规则,Header ID 保持大小写不敏感。更新配置按凭据 ID 的差集删除密文,因此 `Authorization` 改为 `authorization` 不会丢失认证信息。启动时对无歧义的旧环境变量凭据原子迁移密文,不覆盖新 ID 已有的值;若旧配置把 `TOKEN`、`token` 合并存到了同一个 ID,无法推断原来的两个值,会保留旧密文、停用连接并要求重新录入和测试。删除服务器时也会清理这些保留的旧密文。
|
||||
|
||||
### 2.1 JSON 导入与 API Key 填写
|
||||
|
||||
前端支持 NotesAgent 完整/精简配置、单个 `command / args / env` 配置,以及只含一个服务器的 `mcpServers` 包装。JSON 与表单之间切换会补齐数组、对象及超时默认值,并校验字段类型。批量导入暂不支持;后端配置接口仍只接收 NotesAgent DTO,兼容转换发生在前端。
|
||||
|
||||
可以先声明 `secret_environment_keys`,保存后在服务器卡片的密码框填写密钥;也可以把密钥放进 JSON 的 `environment` 或通用配置的 `env`。前端会将已声明的敏感变量,以及名称含 API Key、Token、Secret、Password、Authorization、Cookie、Credential 的常见字段拆出:普通配置请求只包含键名,密钥另经 Secret API 加密保存。其他敏感字段必须显式声明,不能只依赖名称识别;命令与参数中不要携带密钥。
|
||||
|
||||
例如 MiniMax 的输入结构如下,占位值需在自己的本地页面替换,不要把真实密钥贴进聊天或提交到 Git:
|
||||
|
||||
```json
|
||||
{
|
||||
"name": "MiniMax Coding Plan",
|
||||
"command": "uvx",
|
||||
"args": ["--index-url", "https://pypi.tuna.tsinghua.edu.cn/simple", "--with", "mcp<2", "minimax-coding-plan-mcp", "-y"],
|
||||
"environment": {
|
||||
"MINIMAX_API_HOST": "https://api.minimaxi.com",
|
||||
"MINIMAX_API_KEY": "<在本地填入新密钥>"
|
||||
},
|
||||
"secret_environment_keys": ["MINIMAX_API_KEY"],
|
||||
"startup_timeout_seconds": 120,
|
||||
"tool_timeout_seconds": 300
|
||||
}
|
||||
```
|
||||
|
||||
旧版前端将 `environment.MINIMAX_API_KEY` 与 `secret_environment_keys` 原样一起发送,触发后端“普通与敏感变量不可同名”的校验。这是配置保存失败,不是模型服务返回的鉴权失败。现在在前端拆分两类请求,后端仍保留互斥校验。
|
||||
|
||||
导入兼容规则:`env` 转为 `environment`;`timeout` 作为启动超时;`sse_read_timeout` 作为工具等待上限,不保留其原客户端 SSE 读取超时语义。启动超时范围为 1–120 秒,工具超时为 1–300 秒。URL 必须是纯地址,不能使用 Markdown 链接,JSON 中不能包含 `\_` 这样的非法转义。
|
||||
|
||||
另一个已修复的失败原因是运行时适配层复用了 `PluginBackend` 的整数超时与 60 秒启动上限,导致合法的 120 秒或小数超时配置在保存返回、读取或测试时失败。独立 Server 现在使用专门的 Bridge 适配模型,保留自己的浮点超时范围,不改变 Plugin 清单原有约束。已有的 120 秒记录可直接读取,无需删库重建。
|
||||
|
||||
配置保存成功但后续 Secret 写入失败时,窗口保留服务器 ID、新版本和未写入的密钥。点击保存会更新同一服务器并重试,不重复创建记录;取消会清除未保存密钥,已经保存的服务器和凭据不会回滚。错误信息显示在配置窗口内。保存配置不会自动运行第三方进程,仍需确认、测试和启用。
|
||||
|
||||
暂存的 Header Secret 与已保存凭据使用一致的大小写规则:将 `Authorization` 改为 `authorization` 不会丢弃尚未保存的值,提交时采用当前声明名。重新输入同一 Header 的值会覆盖旧草稿;真正删除声明才清除草稿。环境变量仍区分大小写,不会把 `TOKEN` 的草稿转交给 `token`。
|
||||
|
||||
跨 Registry 与凭据存储的删除以“失败后可重试”为顺序约束:先原子清理密文,再提交新版本或删除 Registry 记录。凭据存储失败时保留原版本和 Server 记录,避免出现返回 500 但配置已提交、版本无法重试或密文失去清理入口的状态。
|
||||
|
||||
写入或删除 Secret 会先停用正在运行的连接、注销动态 Tool,并撤销当前配置的测试通过状态;必须使用新凭据重新测试后才能启用。这样页面展示的凭据状态不会与运行中进程实际持有的旧凭据不一致。
|
||||
|
||||
## 3. 启用与运行时规则
|
||||
|
||||
一次连接按以下顺序执行:
|
||||
|
||||
1. 用户检查服务端生成的连接摘要并确认当前摘要;
|
||||
2. 后端临时连接,完成 initialize 和 `tools/list` 后关闭连接;
|
||||
3. 只有当前摘要测试成功,启用操作才会启动长期连接并注册 Tool;
|
||||
4. 停用、删除、超时或异常退出会注销 Tool 并关闭连接。
|
||||
|
||||
运行期连接异常或 Tool 列表变化时,Registry 会在同一生命周期临界区内标记不可用、注销 Tool,并从 Bridge 移除 Host。每次启动分配独立的连接代次;失败回调取得锁后先核对代次,旧连接延迟到达的回调不能停用新连接。停用、测试结束及关闭服务时撤销对应代次。
|
||||
|
||||
所有独立 MCP API 都通过工作线程执行,包括新增、确认授权和读取接口。虽然部分操作不直接访问网络,但仍可能等待正在测试或启动的连接持有的锁,不能在 FastAPI 事件循环上同步等待。
|
||||
|
||||
HTTP Header 中 `Host`、`Content-Type`、`MCP-Session-Id` 等协议保留项不可由配置覆盖。URL 不允许内嵌凭据或 Fragment。旧 SSE 返回的 POST endpoint 必须与初始 URL 同源,防止认证 Header 被转发到其他站点。
|
||||
|
||||
启动与 Tool 请求超时会同时应用于业务等待和底层 HTTP 请求;旧 SSE 的 endpoint 等待也使用启动超时。非主动结束的旧 SSE 事件流视为 Host 不可用,宿主随后注销 Tool。注册表加载时逐条校验 Server ID、配置字段、Transport 组合和摘要格式,损坏记录统一返回 `MCP_REGISTRY_INVALID`。
|
||||
|
||||
两种 HTTP Transport 共用有界 SSE 行解析器:按响应字节块检查未完成行及当前事件的累计大小,再扩展缓冲区,不依赖 `iter_lines()` 先缓存完整行。持续无换行输入也会及时触发上限;解析兼容跨块 UTF-8、首行 BOM、LF/CR/CRLF、多行 data 和事件间计数重置。`tests/test_mcp_sse_limits.py` 覆盖这些边界,防止仅在完整行生成后检查大小。
|
||||
|
||||
stdio 命令不经过 Shell,管道、重定向和命令拼接不会被解释。Windows 使用新进程组并通过 `taskkill /T` 回收子树;POSIX 使用独立 session/process group 并向进程组发信号。Python 阶段仍无法提供文件、网络、系统调用或操作系统版本差异下的绝对隔离保证。
|
||||
|
||||
`uvx` 模板使用 `--isolated`、明确的 `--from` 和固定包版本。它只能隔离依赖,不能替代安全沙箱。非开发环境仍拒绝启动 stdio Server,并返回 `403 MCP_SANDBOX_REQUIRED`;远程 HTTP Transport 不创建本机子进程,但仍要求摘要确认和成功测试。第三阶段前的 C.5 将冻结 Tauri/Rust 沙箱设计。
|
||||
|
||||
## 4. 接口
|
||||
|
||||
```text
|
||||
GET /api/mcp/servers
|
||||
POST /api/mcp/servers
|
||||
GET /api/mcp/servers/{server_id}
|
||||
PUT /api/mcp/servers/{server_id}
|
||||
DELETE /api/mcp/servers/{server_id}
|
||||
GET /api/mcp/servers/{server_id}/tools
|
||||
POST /api/mcp/servers/{server_id}/trust
|
||||
POST /api/mcp/servers/{server_id}/test
|
||||
POST /api/mcp/servers/{server_id}/enable
|
||||
POST /api/mcp/servers/{server_id}/disable
|
||||
PUT /api/mcp/servers/{server_id}/secrets/{key}?kind=environment|header
|
||||
DELETE /api/mcp/servers/{server_id}/secrets/{key}?kind=environment|header
|
||||
```
|
||||
|
||||
完整字段、状态和错误码见《第二阶段接口契约-开发版》。协议实现参考 MCP 官方的 [Transports](https://modelcontextprotocol.io/specification/2025-11-25/basic/transports) 与 [Lifecycle](https://modelcontextprotocol.io/specification/2025-11-25/basic/lifecycle)。
|
||||
|
||||
## 5. 验证
|
||||
|
||||
```powershell
|
||||
cd backend
|
||||
uv run pytest -q tests/test_mcp_registry.py tests/test_extension_core.py
|
||||
|
||||
cd ../frontend
|
||||
npm run type-check
|
||||
npm test
|
||||
npm run build
|
||||
```
|
||||
|
||||
后端测试使用无需网络或真实密钥的 stdio Fixture,以及 `httpx.MockTransport` 驱动的确定性 HTTP/SSE Server Fixture。覆盖摘要授权、Secret 不回显、版本冲突、生产门禁、重启恢复、Streamable HTTP Session/Header/工具调用及旧 SSE 同源校验;新增覆盖旧回调隔离、路由线程卸载、凭据大小写差异与旧密文迁移。前端覆盖模板切换、JSON 默认值与格式兼容、明文拆分、模式切换、部分保存失败重试、取消清理、测试失败和删除确认。此处的 MiniMax 配置转换测试使用假密钥,不等同于真实 MiniMax 网络调用验证。
|
||||
|
||||
## 6. 后续边界
|
||||
|
||||
### 本轮 P1/P2 修复验收
|
||||
|
||||
| 审阅问题 | 修复方式 | 回归验证 |
|
||||
| --- | --- | --- |
|
||||
| P1:新增或授权等待生命周期锁时阻塞事件循环 | 独立 MCP 路由统一交给工作线程 | `test_mcp_lifecycle_lock_contention_keeps_event_loop_responsive`:分别阻塞 create/trust,在锁释放前仍能执行健康检查 |
|
||||
| P2:旧失败回调误停新连接 | 回调在锁内核对连接代次 | `test_old_failure_callback_cannot_stop_replacement_host`:旧回调排队期间重启连接,释放锁后新连接仍可用,当前代次的失败仍正确停用 |
|
||||
| P2:精简 JSON 切换模式或编辑保存报错 | 运行时校验、默认值补全、统一转换 | `configuration.spec.ts` 与 `McpServersView.spec.ts`:精简 JSON、编辑版本、模式切换和 Secret 保存重试 |
|
||||
| P2:Header 大小写改名删除凭据 | 按规范化凭据 ID 而非原始键名计算差集 | `test_header_case_only_rename_preserves_secret` |
|
||||
| P2:大小写不同的环境变量覆盖同一凭据 | 区分大小写的 v2 ID,带迁移标记的旧密文迁移 | `test_environment_secrets_are_case_sensitive_and_delete_independently` 及 legacy migration 测试,包含删除后不复活旧密钥 |
|
||||
| P1:SSE 无换行输入在大小校验前无限缓冲 | 在行拼接前校验字节数与事件累计大小 | `test_mcp_sse_limits.py`,包括小块持续输入和跨块换行 |
|
||||
| P2:Header 大小写改名丢失未保存密钥 | 草稿使用规范化名称匹配,并重新绑定当前声明名 | `configuration.spec.ts` 与页面保存回归测试 |
|
||||
|
||||
这些修复不放宽 stdio 的 JSON-RPC 校验。第三方程序向 stdout 打印普通日志造成的握手失败,应由服务端调整输出或使用不打印日志的启动入口处理。
|
||||
|
||||
### 后续工作
|
||||
|
||||
- 增加真实第三方 Server 的兼容矩阵;确定性 Fixture 只能证明宿主协议行为,不能代表所有实现兼容;
|
||||
- C.5 在第二阶段开发与测试完成后、第三阶段桌面端实现前冻结沙箱 Contract;
|
||||
- 第三阶段将 stdio 进程创建和 Secret 托管迁移至 Tauri/Rust Host 与 Stronghold/系统 Keychain。
|
||||
@@ -1,9 +1,11 @@
|
||||
# 第一阶段测试验证操作手册
|
||||
|
||||
> 适用基线:2026-08-30 `main`
|
||||
> 适用基线:2026-09-02 `main`
|
||||
> 适用对象:开发、自测、代码审阅、合并验收和 Demo 前检查
|
||||
> 验证范围:Vue Web 前端、FastAPI、Knowledge/Retrieval Core、AI/Agent Core、Extension Core、Provider 与开发阶段凭据链路
|
||||
|
||||
> 状态说明:本文保留第一阶段功能验收口径;当前全量自动化回归同时覆盖第二阶段已经合并的 Agent Trace 持久化与 SSE 恢复、stdio MCP Host、Plugin Command/Settings Contract 和前端 Service,但不反向修改第一阶段的交付定义。
|
||||
|
||||
## 1. 验证目标
|
||||
|
||||
本手册用于确认第一阶段已经形成可运行的本地知识工作流:
|
||||
@@ -18,7 +20,7 @@
|
||||
→ Skill / Plugin 完成生命周期与 Tool 注册
|
||||
```
|
||||
|
||||
当前不作为第一阶段通过条件的内容:Tauri/Rust Host、Stronghold、真实桌面文件系统、独立 MCP Plugin Host、真实音频模型和 Sync Server。
|
||||
当前不作为第一阶段通过条件的内容:Tauri/Rust Host、Stronghold、原生多 Vault 文件系统、真实音频模型和 Sync Server。独立 stdio MCP Plugin Host 已在第二阶段完成,并进入当前全量回归测试。
|
||||
|
||||
## 2. 环境准备
|
||||
|
||||
@@ -68,7 +70,7 @@ uv run pytest -q -p no:cacheprovider
|
||||
当前基线:
|
||||
|
||||
```text
|
||||
81 passed
|
||||
136 passed
|
||||
```
|
||||
|
||||
通过标准:退出码为 0、失败数为 0。用例数可以随功能增加,但不得低于当前基线。
|
||||
@@ -83,11 +85,11 @@ pnpm test
|
||||
当前基线:
|
||||
|
||||
```text
|
||||
11 test files passed
|
||||
27 tests passed
|
||||
12 test files passed
|
||||
29 tests passed
|
||||
```
|
||||
|
||||
通过标准:退出码为 0、失败数为 0。测试覆盖 Provider Store、主题偏好、Workspace、文件树、文件切换、可视化编辑器、智能体中文标签、轻量动效性能约束、Markdown 对比度 Token、scoped CSS 选择器约束和 Shiki GitHub 双主题输出。
|
||||
通过标准:退出码为 0、失败数为 0。测试覆盖 Provider Store、主题偏好、Workspace、文件树、文件切换、可视化编辑器、智能体中文标签、SSE 恢复游标、Plugin Command/Settings Service、轻量动效性能约束、Markdown 对比度 Token、scoped CSS 选择器约束和 Shiki GitHub 双主题输出。
|
||||
|
||||
### 3.3 类型检查与生产构建
|
||||
|
||||
@@ -330,7 +332,7 @@ Invoke-RestMethod -Method Delete -Uri "$apiBase/notes/$noteId"
|
||||
- 没有真实密钥、生成目录或运行数据进入 Git;
|
||||
- 已记录测试环境、提交、结果、警告和遗留问题。
|
||||
|
||||
外部 OpenAI/DeepSeek、Tauri、Stronghold、真实文件系统、真实音频模型与 Sync Server 失败或未测,不阻止当前第一阶段 Web 联调基线通过,但必须在验收记录中注明“未纳入本阶段”或“选测未执行”。
|
||||
外部 OpenAI/DeepSeek、Tauri、Stronghold、原生多 Vault 文件系统、真实音频模型与 Sync Server 失败或未测,不阻止当前第一阶段 Web 联调基线通过,但必须在验收记录中注明“未纳入本阶段”或“选测未执行”。
|
||||
|
||||
## 9. 验收记录模板
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
> 本文记录 `feat/knowledge-retrieval-core` 合并前后的两轮代码审阅、问题复现、修复过程与工程经验。
|
||||
> 它既是团队内部的问题档案,也可作为后续技术文档、课程报告和博客文章的素材底稿。
|
||||
|
||||
> 2026-08-30 状态补充:本文中的 43 项测试是当时该模块的历史基线,不应替换为当前全仓测试数。相关修复仍有效,当前完整后端回归为 71 项通过,Knowledge/Retrieval 已通过 `notes.*` 与 `rag.search` Tool 接入 Agent Runtime。
|
||||
> 2026-09-02 状态补充:本文中的 43 项测试是当时该模块的历史基线,不应替换为当前全仓测试数。相关修复仍有效,当前完整后端回归为 136 项通过,Knowledge/Retrieval 已通过 `notes.*` 与 `rag.search` Tool 接入 Agent Runtime。
|
||||
|
||||
## 1. 背景
|
||||
|
||||
|
||||
@@ -0,0 +1,464 @@
|
||||
# Plugin Command 与 Settings 问题与修复复盘
|
||||
|
||||
> 审阅与修复日期:2026-09-02
|
||||
> 涉及分支:`feat/plugin-command-settings`
|
||||
> 初始功能提交:`6a08ad8 feat(extension): 实现插件命令与设置贡献`
|
||||
> 最终修复提交:`d39ae72 fix(extension): 收紧插件命令运行时契约`
|
||||
> 合并提交:`ff3da5d Merge pull request 'feat(extension): 实现 Plugin Command 与 Settings Contribution' (#10)`
|
||||
> 文档用途:记录阶段 D 从首次提交、连续审阅到合并期间发现的问题,说明形成原因、实际后果、解决思路和最终方案,供后续开发、问题定位、比赛材料和技术博客写作使用。
|
||||
|
||||
## 1. 背景与结论
|
||||
|
||||
阶段 D 在既有 Extension Core、MCP Bridge 和 Plugin Host 上增加 Plugin Command 与 Settings Contribution。初始实现已经具备命令注册、参数校验、设置持久化、Secret 加密存储和前端 Service,但首次提交仍把若干“接口能够调用”误当成了“安全边界已经闭合”。
|
||||
|
||||
本次分支没有在初始功能提交后直接合并,而是围绕清单输入、凭据隔离、JSON Schema、MCP 协议、事务一致性和前后端 Contract 连续审阅。最终共形成 1 个功能提交和 6 个修复提交,处理 12 类问题:
|
||||
|
||||
| 编号 | 问题 | 级别 | 处理结果 |
|
||||
| --- | --- | --- | --- |
|
||||
| D-01 | Plugin Secret 可被通用凭据接口或 Provider 引用 | P0 | 建立 `plugin.*` 保留命名空间和双层访问门禁 |
|
||||
| D-02 | Command 与 Settings 清单边界不完整 | P1 | 收紧空值、重复项、权限、类型和数值边界校验 |
|
||||
| D-03 | MCP Command Target 只有声明,没有真实执行链路 | P0 | 增加专用 MCP 命令目标并与 Agent Tool 隔离 |
|
||||
| D-04 | JSON Schema 外部引用可能触发越权 I/O | P0 | 只允许当前文档内 Fragment 引用 |
|
||||
| D-05 | Secret 引用可被篡改并跨凭据命名空间 | P0 | 改用定长哈希引用并在读取、删除前重新核对 |
|
||||
| D-06 | Secret 删除和插件卸载缺少事务一致性 | P0 | 原子替换凭据表,失败时恢复 Settings 引用 |
|
||||
| D-07 | Schema 引用校验忽略嵌套 `$id` 资源作用域 | P1 | 使用 Draft 2020-12 Resource Resolver 按资源解析 |
|
||||
| D-08 | MCP Command 未收到普通 Settings | P1 | 在固定信封中同时传递 Settings 与声明过的 Secret |
|
||||
| D-09 | Command Effect 是任意字典,前后端边界过宽 | P0 | 改为五类判别联合并限制各自 Payload |
|
||||
| D-10 | 必填普通设置没有阻止启用和执行 | P1 | 在 Enable 与 Execute 两个入口增加运行时门禁 |
|
||||
| D-11 | MCP 工具 Schema 被发现后丢弃,真实信封未校验 | P0 | 保存目标 Schema,并在调用前校验完整信封 |
|
||||
| D-12 | 启用探测和回归测试存在假阳性或假阴性 | P1 | 使用最小协议标记,移除自制求解器并补强测试 |
|
||||
|
||||
最终验证基线为后端 136 项测试、前端 29 项测试、TypeScript 类型检查、前端生产构建、`uv lock --check` 和 `git diff --check` 通过。PR #10 已合并到 `main`。
|
||||
|
||||
## 2. 提交与审阅过程
|
||||
|
||||
| 顺序 | 提交 | 主要内容 |
|
||||
| --- | --- | --- |
|
||||
| 1 | `6a08ad8` | 首次实现 Plugin Command、Settings、Secret API 和前端 Service |
|
||||
| 2 | `c3ef9df` | 收紧 Plugin Secret、Provider 凭据和 Command 清单边界 |
|
||||
| 3 | `9e680a0` | 补齐真实 MCP Command Target,限制 JSON Schema 外部引用 |
|
||||
| 4 | `022c322` | 改造 Secret 引用,修复删除和卸载事务 |
|
||||
| 5 | `c06b962` | 按 Draft 2020-12 资源作用域修复 Schema 引用解析 |
|
||||
| 6 | `eb3464b` | 补回 MCP Command 信封中的普通 Settings |
|
||||
| 7 | `d39ae72` | 收紧 Effect、必填设置和 MCP 运行时 Contract,修复测试盲点 |
|
||||
| 合并 | `ff3da5d` | PR #10 合并进入 `main` |
|
||||
|
||||
这段过程说明,Extension Core 的风险主要不在正常路径能否运行,而在不同入口是否共享同一套边界。HTTP、Provider Factory、Plugin Runtime、MCP Host、Settings Store 和测试 Fixture 只要有一个入口绕过限制,就可能形成跨命名空间访问、错误执行或假测试通过。
|
||||
|
||||
## 3. D-01:Plugin Secret 可被通用凭据接口或 Provider 引用
|
||||
|
||||
### 原因
|
||||
|
||||
初始实现把 Plugin Secret 存入已有 `EncryptedCredentialStore`,但只在 Plugin Settings API 中隐藏明文,没有为凭据 ID 建立用途隔离。通用凭据 API 可以读写或删除同名 ID,Provider 配置也可以把 Plugin Secret 的凭据引用当作自己的 API Key 使用。
|
||||
|
||||
此外,Command 能否读取 Secret 只依赖 Plugin 拥有 `secrets.use` 权限,没有继续限制到当前 Command 在 `commands.yaml` 中明确声明的字段。
|
||||
|
||||
### 后果
|
||||
|
||||
- 前端或其他模块可以覆盖、删除 Plugin 私有 Secret;
|
||||
- Provider 可能把插件密钥发送给外部模型服务;
|
||||
- 同一 Plugin 内权限较低的 Command 可以读取与自己无关的 Secret;
|
||||
- API 没有回显明文并不代表 Secret 没有发生横向流动。
|
||||
|
||||
### 解决思路
|
||||
|
||||
Secret 隔离必须同时覆盖存储命名空间、公共 HTTP 入口、Provider 解析入口和 Command 字段级授权,不能只在返回 DTO 上隐藏值。
|
||||
|
||||
### 解决方案
|
||||
|
||||
- 将 `plugin.*` 设为 Plugin Settings 专用保留命名空间;
|
||||
- 通用凭据查询、写入和删除接口拒绝访问该前缀;
|
||||
- 在 `ProviderFactory` 外包一层 `ProviderCredentialResolver`,即使绕过 HTTP 配置校验也不能读取 Plugin Secret;
|
||||
- Command 通过 `secrets` 字段声明允许读取的 Setting Key;
|
||||
- Resolver 同时检查插件权限、字段是否存在、字段类型和 Command 声明;
|
||||
- 未声明字段返回 `PLUGIN_SECRET_ACCESS_DENIED`,必填 Secret 未配置返回 `PLUGIN_SECRET_REQUIRED`;
|
||||
- 审计事件不记录 arguments、effect 或 Secret。
|
||||
|
||||
## 4. D-02:Command 与 Settings 清单边界不完整
|
||||
|
||||
### 原因
|
||||
|
||||
YAML 解析成功只说明语法可读,并不代表清单结构满足宿主协议。初始校验对 `commands:` 空值、重复 Secret、未声明权限、无效执行目标、Settings 默认值及非有限数值边界等情况覆盖不足。
|
||||
|
||||
例如 YAML 中只有 `commands:` 时,解析结果是 `null` 而不是空数组;`NaN` 和正负无穷虽然属于 Python 浮点值,却不能成为可移植的表单边界。
|
||||
|
||||
### 后果
|
||||
|
||||
- 安装阶段可能放过无法执行的 Command;
|
||||
- 空清单在后续遍历时变成内部异常,而不是稳定业务错误;
|
||||
- 前后端对数值范围和默认值产生不一致理解;
|
||||
- 重复声明或越过 Plugin 命名空间的 ID 会污染全局注册表;
|
||||
- 错误只能到运行期暴露,定位成本更高。
|
||||
|
||||
### 解决思路
|
||||
|
||||
把 Plugin 包视为不可信输入,在安装阶段完成结构、语义和权限的完整验证,并将解析异常统一转换成稳定的 `ExtensionError`。
|
||||
|
||||
### 解决方案
|
||||
|
||||
- 要求 `contributes` 声明与 `commands.yaml`、`settings.yaml` 内容完全一致;
|
||||
- 拒绝 `null` Command 列表、重复 ID、重复 Secret 和越过 Plugin 命名空间的标识;
|
||||
- 执行目标只能在受控 `handler` 与当前 Plugin 的 `mcp_tool` 中二选一;
|
||||
- 校验 `when`、Context、Location 和权限白名单;
|
||||
- Settings 类型限定为 `string`、`number`、`boolean`、`select`、`secret`;
|
||||
- 校验默认值、Select 选项、必填规则和最小/最大边界;
|
||||
- 拒绝 `NaN`、正无穷和负无穷;
|
||||
- 清单错误统一返回可定位的稳定错误码。
|
||||
|
||||
## 5. D-03:MCP Command Target 只有声明,没有真实执行链路
|
||||
|
||||
### 原因
|
||||
|
||||
初始 Command 执行器只支持宿主内置 Handler。接口和规划中虽然存在 MCP Command 的概念,但 Runtime 没有把 Command 绑定到 MCP 工具,也没有定义 Command Context、Settings 和 Secret 如何进入 MCP 请求。
|
||||
|
||||
直接复用 Agent `ToolRegistry` 看似省事,却会把“用户主动执行的插件命令”和“模型可自主调用的 Agent Tool”混成同一种能力。
|
||||
|
||||
### 后果
|
||||
|
||||
- MCP 插件声明的 Command 实际无法执行;
|
||||
- 如果简单注册为 Agent Tool,模型可能绕过命令位置、Context 裁剪和 Secret 声明直接调用;
|
||||
- Command 超时、Effect 和错误边界无法统一;
|
||||
- 前端看到命令已注册,点击后却只能得到运行时错误。
|
||||
|
||||
### 解决思路
|
||||
|
||||
为 MCP Command 建立独立执行通道。它可以复用 MCP 连接,但不能自动进入 Agent Tool Registry;宿主负责构造固定协议信封并验证返回的白名单 Effect。
|
||||
|
||||
### 解决方案
|
||||
|
||||
- `commands.yaml` 支持当前 Plugin 命名空间内的 `mcp_tool` 目标;
|
||||
- Runtime 在 Plugin 启用后绑定目标,在禁用、异常退出和重启时同步注销;
|
||||
- MCP Command 不注册到 Agent `ToolRegistry`;
|
||||
- 宿主只传递已校验 arguments、已裁剪 Context、有效 Settings 和当前 Command 声明的 Secret;
|
||||
- MCP 返回结果必须转换成受控 Effect;
|
||||
- 远程原始异常、过大结果和超时统一映射为稳定宿主错误。
|
||||
|
||||
## 6. D-04:JSON Schema 外部引用可能触发越权 I/O
|
||||
|
||||
### 原因
|
||||
|
||||
Command 和 MCP Tool 都接受插件提供的 JSON Schema。初始实现直接交给校验器处理 `$ref` 或 `$dynamicRef`,没有限制引用 URI。恶意或错误 Schema 可以引用本地文件、HTTP 地址或其他外部资源。
|
||||
|
||||
### 后果
|
||||
|
||||
- Schema 校验可能读取宿主文件或发起未授权网络请求;
|
||||
- 安装一个插件就可能产生隐式 I/O;
|
||||
- 离线环境中校验结果不稳定;
|
||||
- 外部资源变化会让相同插件包得到不同验证结果;
|
||||
- Command 和 Agent Tool 如果采用不同规则,会出现新的绕过路径。
|
||||
|
||||
### 解决思路
|
||||
|
||||
当前阶段不需要跨文件 Schema。宿主应只允许当前文档内部的 Fragment 引用,并在注册阶段递归检查所有 Schema 节点。
|
||||
|
||||
### 解决方案
|
||||
|
||||
- Command 与 Tool 共用 `schema_security` 校验边界;
|
||||
- 递归扫描 `$ref` 和 `$dynamicRef`;
|
||||
- 只允许以 `#` 开头的当前文档 Fragment;
|
||||
- 拒绝文件、HTTP 和其他外部资源 URI;
|
||||
- 无法解析的本地引用在安装或注册阶段直接失败;
|
||||
- 运行期继续使用官方 Draft 2020-12 Validator 校验数据。
|
||||
|
||||
## 7. D-05:Secret 引用可被篡改并跨凭据命名空间
|
||||
|
||||
### 原因
|
||||
|
||||
初始 Settings 文件保存的 Secret 引用由可读的 Plugin ID 和 Setting Key 拼接而成。长度随名称增长,而且 Runtime 读取引用时默认信任磁盘内容,没有重新确认该引用确实属于当前字段。
|
||||
|
||||
本地文件损坏或被篡改后,一个 Plugin Setting 可以被改为指向 Provider 凭据或另一个 Plugin 的 Secret。
|
||||
|
||||
### 后果
|
||||
|
||||
- 长 Plugin ID 和 Setting Key 可能超过凭据 ID 长度限制;
|
||||
- Settings 文件泄露内部字段名称;
|
||||
- 篡改引用可能造成跨命名空间读取或删除;
|
||||
- 卸载一个插件时可能误删其他模块的凭据。
|
||||
|
||||
### 解决思路
|
||||
|
||||
Secret 引用应由宿主确定性生成,长度固定,并在每次敏感操作前由当前 `plugin_id + setting_key` 重新计算,而不是信任持久化文件。
|
||||
|
||||
### 解决方案
|
||||
|
||||
使用以下语义生成引用:
|
||||
|
||||
```text
|
||||
plugin.<sha256(plugin_id + "\\0" + setting_key)>
|
||||
```
|
||||
|
||||
- 引用长度固定且符合凭据 ID 规则;
|
||||
- Settings Store 只保存引用和 `configured` 状态,不保存明文;
|
||||
- 读取、覆盖、删除和卸载前重新计算期望引用;
|
||||
- 引用不匹配时按损坏存储拒绝处理;
|
||||
- 增加跨命名空间篡改回归测试。
|
||||
|
||||
## 8. D-06:Secret 删除和插件卸载缺少事务一致性
|
||||
|
||||
### 原因
|
||||
|
||||
Secret 同时涉及 Settings 引用文件和加密凭据文件。初始删除流程按顺序修改两个存储,但任一步失败都没有完整回滚。插件卸载多个 Secret 时逐条删除,执行到一半失败会留下部分清理状态。
|
||||
|
||||
### 后果
|
||||
|
||||
- Settings 显示未配置,但密文仍残留;
|
||||
- 凭据已删除,Settings 却仍显示已配置;
|
||||
- 多 Secret 卸载可能只删除前几项;
|
||||
- 重试操作无法判断上一次执行到哪里;
|
||||
- 用户以为插件卸载已清除密钥,实际磁盘仍可能保留数据。
|
||||
|
||||
### 解决思路
|
||||
|
||||
把引用更新和凭据删除看作一个逻辑事务。底层单文件凭据表应一次构造新状态并原子替换;跨 Settings 与 Credential Store 的操作需要显式补偿回滚。
|
||||
|
||||
### 解决方案
|
||||
|
||||
- `EncryptedCredentialStore` 增加多凭据原子删除;
|
||||
- 先验证全部目标引用,再生成新的凭据表;
|
||||
- 通过临时文件和原子替换一次提交;
|
||||
- 删除单个 Secret 时,凭据删除失败则恢复原 Settings 引用;
|
||||
- 卸载 Plugin 时,批量删除失败则恢复完整 Settings 命名空间;
|
||||
- 错误统一转换为稳定存储错误,避免部分成功被当作完整成功。
|
||||
|
||||
## 9. D-07:Schema 引用校验忽略嵌套 `$id` 资源作用域
|
||||
|
||||
### 原因
|
||||
|
||||
第一版本地引用检查把整个 Schema 当成单一 Fragment 树,用根文档指针或 Anchor 查找所有引用。Draft 2020-12 允许嵌套 `$id` 创建新的 Schema Resource;资源内部的 `#anchor` 应相对于新的 Base URI 解析,根资源也不能反向使用嵌套资源的 Anchor。
|
||||
|
||||
### 后果
|
||||
|
||||
- 合法的嵌套资源引用被误拒绝;
|
||||
- 根 Schema 可能错误引用只属于子资源的 Anchor;
|
||||
- 宿主预检结果与官方运行时 Validator 不一致;
|
||||
- 同一 Schema 在安装阶段通过,却可能在执行阶段失败,反之亦然。
|
||||
|
||||
### 解决思路
|
||||
|
||||
安全限制仍然是“禁止外部资源”,但本地资源内部的解析语义必须遵守 JSON Schema 标准,不能自己用字符串和全局 Anchor 集合近似实现。
|
||||
|
||||
### 解决方案
|
||||
|
||||
- 引入与 Draft 2020-12 Validator 配套的 Resource Registry;
|
||||
- 为根资源和嵌套 `$id` 建立正确作用域;
|
||||
- 每个引用按其所在资源的 Base URI 解析;
|
||||
- 保留外部资源拒绝策略;
|
||||
- 增加“根资源不能使用子资源 Anchor”和“子资源可使用自身 Anchor”的成对测试。
|
||||
|
||||
## 10. D-08:MCP Command 未收到普通 Settings
|
||||
|
||||
### 原因
|
||||
|
||||
真实 MCP Command 链路补齐后,固定信封传递了 Command ID、Arguments、Context 和 Secret,但遗漏了已经通过 Schema 校验的普通 Settings。声明式内置 Handler 能读取 Settings,MCP Handler 却不能,两个执行目标语义不一致。
|
||||
|
||||
### 后果
|
||||
|
||||
- 用户在插件设置页修改普通配置,对 MCP Command 不生效;
|
||||
- 插件只能把非敏感设置错误地编码进 arguments 或 Secret;
|
||||
- 内置 Handler 测试通过会掩盖 MCP 路径的缺口;
|
||||
- 插件从内置实现迁移到 MCP 后行为发生变化。
|
||||
|
||||
### 解决思路
|
||||
|
||||
内置 Handler 与 MCP Command 应消费同一份运行时配置。差别只在执行介质,不在 Command Contract。
|
||||
|
||||
### 解决方案
|
||||
|
||||
- MCP 固定信封增加 `settings`;
|
||||
- Settings 由 `PluginSettingsStore.runtime_values()` 产生;
|
||||
- 只传递当前 Plugin Schema 中有效的非敏感字段;
|
||||
- Secret 继续放在独立 `secrets` 命名空间;
|
||||
- Fixture 回显非敏感配置用于断言,但不回显 Secret;
|
||||
- 增加 Settings 实际到达 MCP Server 的集成测试。
|
||||
|
||||
## 11. D-09:Command Effect 是任意字典,前后端边界过宽
|
||||
|
||||
### 原因
|
||||
|
||||
初始 `PluginCommandEffect` 只有 `type` 和任意 `payload`。宿主虽然限制 Effect 名称和总体大小,却没有限制 Payload 字段、路由名称、刷新范围或通知级别。前端 TypeScript 也只能把 Payload 当成普通对象处理。
|
||||
|
||||
### 后果
|
||||
|
||||
- 插件可以返回前端从未支持的字段和路由;
|
||||
- 前端需要在运行时猜测 Payload 结构;
|
||||
- `navigate` 或 `refresh` 可能越过宿主允许的目标;
|
||||
- OpenAPI 无法表达不同 Effect 的必填字段;
|
||||
- 无效 Effect 往往要到页面执行时才暴露。
|
||||
|
||||
### 解决思路
|
||||
|
||||
Effect 是宿主能力协议,不是插件任意消息。每种 Effect 都应是独立、封闭、可判别的 Contract,并在进入 HTTP 响应前完成验证。
|
||||
|
||||
### 解决方案
|
||||
|
||||
- 后端建立 `notification`、`navigate`、`refresh`、`job`、`none` 五类模型;
|
||||
- 使用 `type` 作为 Pydantic 判别字段;
|
||||
- 限制通知级别、消息长度、路由白名单、刷新 Scope 和 Job ID;
|
||||
- TypeScript 同步为精确的判别联合;
|
||||
- OpenAPI `PluginCommandResult.effect` 生成 `oneOf` 和 discriminator;
|
||||
- 保留可序列化性和 64 KiB 总大小限制作为第二层保护。
|
||||
|
||||
## 12. D-10:必填普通设置没有阻止启用和执行
|
||||
|
||||
### 原因
|
||||
|
||||
初始 Runtime 会合并已保存值和默认值,却没有检查 `required` 且没有默认值的普通字段是否仍为空。Secret 已有独立缺失检查,因此测试容易只覆盖 Secret,忽略普通 Settings。
|
||||
|
||||
### 后果
|
||||
|
||||
- 配置不完整的 Plugin 仍可启动 MCP Host;
|
||||
- Command 到插件内部才因缺少字段失败;
|
||||
- 用户只能看到模糊执行错误,不知道应先补配置;
|
||||
- 插件启用后删除必填值,后续执行没有再次校验。
|
||||
|
||||
### 解决思路
|
||||
|
||||
必填设置既是启用前置条件,也是每次执行的运行时不变量。不能只在保存表单时校验,因为磁盘内容可能变化,启用后的配置也可能被更新。
|
||||
|
||||
### 解决方案
|
||||
|
||||
- `runtime_values()` 返回完整有效普通 Settings;
|
||||
- 缺少必填且无默认值的字段时抛出 `PLUGIN_SETTINGS_REQUIRED`;
|
||||
- Plugin Enable 前执行一次检查,不启动无效 Host;
|
||||
- 每次 Command Execute 前重新检查;
|
||||
- 返回 409 和缺失字段上下文,便于前端引导用户进入设置页;
|
||||
- 增加“补齐设置后可启用”的完整回归测试。
|
||||
|
||||
## 13. D-11:MCP 工具 Schema 被发现后丢弃,真实信封未校验
|
||||
|
||||
### 原因
|
||||
|
||||
MCP 初始化阶段能够取得工具名称和 `inputSchema`,但 Runtime 记录只保留了工具名。Command 执行时直接发送宿主信封,没有用目标工具的完整 Schema 校验实际数据。
|
||||
|
||||
这意味着插件只要暴露同名工具就可能通过启用检查,即使它根本不接受 NotesAgent Command 协议。
|
||||
|
||||
### 后果
|
||||
|
||||
- 不兼容目标在启用阶段被注册为可执行 Command;
|
||||
- 错误推迟到远程 MCP Server,返回信息不稳定;
|
||||
- Context、Settings 或 Secret 结构变化时无法在宿主边界发现漂移;
|
||||
- 前端看到可用命令,执行后才得到 502;
|
||||
- 禁用或重启后若 Schema 缓存不清理,还可能使用过期契约。
|
||||
|
||||
### 解决思路
|
||||
|
||||
发现阶段保留完整目标 Schema;启用阶段只检查最低协议标记;执行阶段再用真实数据验证全部约束。Schema 生命周期必须和 MCP 工具生命周期一致。
|
||||
|
||||
### 解决方案
|
||||
|
||||
- Plugin 运行记录增加 MCP Command Schema 映射;
|
||||
- 禁用、回滚、重启和 Host 不可用时同步清理 Schema;
|
||||
- 启用时要求目标 Schema 顶层直接声明 `properties._notesagent`,且类型为 `object`;
|
||||
- 执行前构造完整固定信封;
|
||||
- 使用官方 Draft 2020-12 Validator 校验真实信封;
|
||||
- 不匹配时返回 `PLUGIN_COMMAND_TARGET_SCHEMA_MISMATCH`,不调用 MCP Server。
|
||||
|
||||
## 14. D-12:启用探测和回归测试存在假阳性或假阴性
|
||||
|
||||
### 原因
|
||||
|
||||
审阅期间先后暴露了三类测试方法问题:
|
||||
|
||||
1. 启用阶段曾用空 arguments、Context、Settings 和 Secret 伪造信封,以判断目标 Schema 是否兼容。合法 Schema 如果要求真实业务字段,会被错误拒绝;
|
||||
2. 为避免空数据误判,曾尝试加入自定义 JSON Schema 可满足性求解,但该近似实现无法正确覆盖 `not`、`oneOf` 等完整 Draft 2020-12 语义;
|
||||
3. MCP Secret Fixture 即使没有收到 `api_key` 也会返回成功,测试只验证了命令成功,没有证明 Secret 真正到达服务端。
|
||||
|
||||
最后还发现空 Echo 消息会构造 `notification`,但新的通知 Contract 要求消息非空,导致合法空输入被包装为 502。
|
||||
|
||||
### 后果
|
||||
|
||||
- 合法插件可能在启用阶段被拒绝,形成假阴性;
|
||||
- 不兼容 Schema 可能被自定义求解器放过,形成假阳性;
|
||||
- Secret 传递链路回归后测试仍显示通过;
|
||||
- Contract 收紧后,旧 Fixture 的边界值会产生新的运行时错误;
|
||||
- 测试数量增加,却没有真正覆盖需要证明的安全事实。
|
||||
|
||||
### 解决思路
|
||||
|
||||
启用阶段只做稳定且明确的协议结构检查,完整 Schema 语义交给官方 Validator 和真实执行数据。测试必须让目标事实缺失时明确失败,而不是通过返回值间接猜测。
|
||||
|
||||
### 解决方案
|
||||
|
||||
- 取消空业务数据 Probe;
|
||||
- 将启用门槛缩小为直接 `_notesagent: { type: object }` 协议标记;
|
||||
- 允许 `$ref`、`$dynamicRef`、`allOf`、`anyOf`、`oneOf` 和 `not` 等约束出现在 `_notesagent` 内部;
|
||||
- 完全移除自定义 Schema 可满足性求解器;
|
||||
- 实际执行统一交给官方 Draft 2020-12 Validator;
|
||||
- MCP Fixture 未收到声明的 Secret 时主动返回 MCP 错误,但永不返回明文;
|
||||
- 空 Echo 返回 `none` Effect,非空 Echo 返回 `notification`;
|
||||
- 为协议标记、真实信封、Secret 到达、空 Echo 和恶意 Effect 分别增加回归测试。
|
||||
|
||||
## 15. 验证方法与结果
|
||||
|
||||
本分支最终执行以下验证:
|
||||
|
||||
```powershell
|
||||
cd backend
|
||||
uv lock --check
|
||||
uv run pytest
|
||||
|
||||
cd ../frontend
|
||||
pnpm test -- --run
|
||||
pnpm type-check
|
||||
pnpm build
|
||||
|
||||
cd ..
|
||||
git diff --check
|
||||
```
|
||||
|
||||
结果:
|
||||
|
||||
```text
|
||||
backend: 136 passed
|
||||
frontend: 29 passed
|
||||
TypeScript type-check passed
|
||||
frontend production build passed
|
||||
uv lock --check passed
|
||||
git diff --check passed
|
||||
```
|
||||
|
||||
后端测试在 Windows 下仍有既有 `.pytest_cache` 权限警告,不影响测试结果。前端生产构建仍有既有大 Chunk 提示,不影响本次功能正确性和合并结论。
|
||||
|
||||
新增或强化的关键测试包括:
|
||||
|
||||
- Plugin Secret 公共 API、Provider Resolver 和 Command Resolver 三层隔离;
|
||||
- 清单空值、重复项、无效权限、非有限边界和非法 Schema;
|
||||
- MCP Command 与 Agent Tool 隔离;
|
||||
- 固定信封中的 Context、Settings 和声明式 Secret;
|
||||
- 外部 Schema 引用拒绝和嵌套 `$id` 资源作用域;
|
||||
- Secret 定长引用、篡改阻断、删除回滚和卸载原子性;
|
||||
- 五类 Effect 的合法 Payload 和恶意 Payload 拒绝;
|
||||
- 必填普通设置对 Enable 和 Execute 的双重门禁;
|
||||
- MCP 实际信封校验、Secret 到达证明和空 Echo 行为。
|
||||
|
||||
## 16. 预防措施
|
||||
|
||||
- Plugin 包、MCP Schema 和 MCP 返回值一律按不可信输入处理;
|
||||
- Secret 安全检查必须覆盖写入、读取、删除、解析、传输和审计全链路;
|
||||
- 保留命名空间需要在公共 API 和内部 Resolver 两侧同时阻断;
|
||||
- 多存储更新必须设计原子提交或补偿回滚,并测试中途失败;
|
||||
- 不自行实现通用 JSON Schema 求解器,标准语义交给官方 Validator;
|
||||
- 启用阶段只检查静态协议能力,不使用伪业务数据推导可执行性;
|
||||
- Agent Tool 和用户触发的 Plugin Command 必须保持独立注册和权限边界;
|
||||
- Pydantic Contract、OpenAPI、TypeScript DTO 和测试 Fixture 在同一提交中同步;
|
||||
- 回归测试应直接证明目标事实,例如“Secret 确实到达且未落盘”,不能只证明接口返回成功;
|
||||
- 每次收紧 Contract 后,重新检查空值、默认值、最大值和旧 Fixture 等边界输入。
|
||||
|
||||
## 17. 当前边界与后续事项
|
||||
|
||||
本次合并完成的是阶段 D 的宿主协议和开发期运行边界,不等于第三方插件已经具备生产级操作系统隔离。
|
||||
|
||||
当前仍保留以下后续事项:
|
||||
|
||||
- 阶段 C.5 按规划放在第二阶段功能与测试完成后、第三阶段桌面集成正式构建 Tauri/Rust 沙箱之前;
|
||||
- 生产环境继续通过配置门禁拒绝未沙箱化的 MCP Host;
|
||||
- 桌面端将 Plugin Secret 从开发期 Fernet 文件迁移到 Stronghold 或系统 Keychain,保持现有引用和 HTTP Contract;
|
||||
- 前端已实现命令面板、Plugin 详情命令和动态 Settings / Secret 表单;上下文菜单与 Toolbar 挂载点仍复用现有 Service 继续扩展;
|
||||
- Command 审计当前是 500 条有界内存队列,长期审计持久化需在后续阶段单独设计;
|
||||
- 前端大 Chunk 应通过路由和 Markdown 依赖拆包处理,不与本次 Extension Contract 修改混合。
|
||||
|
||||
## 18. 复盘结论
|
||||
|
||||
本次最重要的经验是:插件系统的正确性不能只按“命令是否执行成功”判断。真正需要审阅的是数据从清单进入注册表、从 Settings 进入运行时、从 Secret Store 进入执行器、从宿主信封进入 MCP,以及从 MCP Effect 返回前端的每一道边界。
|
||||
|
||||
连续审阅避免了跨命名空间 Secret 访问、非事务删除、外部 Schema I/O、MCP 伪兼容和任意 Effect 等问题进入 `main`。最终实现把每个边界落到明确 Contract、稳定错误码和可失败的回归测试上,为下一阶段前端集成和桌面安全沙箱提供了可复用基础。
|
||||
@@ -4,7 +4,7 @@
|
||||
> 涉及提交:`f9efc4f`,合并提交 `c6c28e4`。
|
||||
> 文档用途:记录前端分支合并后暴露的问题域、形成原因、实际后果、修复思路和落地方案,供后续技术文档、比赛材料与博客写作使用。
|
||||
|
||||
> 2026-08-30 状态补充:在本文两轮修复之后,项目又完成 Milkdown 写作工具栏、文件切换二次竞态修复、Shiki 只读高亮、Provider 预设/模型发现/加密 API Key 输入,以及智能体页面汉化。当前前端回归基线为 14 项测试和生产构建通过。
|
||||
> 2026-09-02 状态补充:在本文多轮修复之后,项目又完成 Milkdown 写作工具栏、文件切换二次竞态修复、Shiki 只读高亮、Provider 预设/模型发现/加密 API Key 输入、智能体页面汉化,以及 Plugin Command/Settings Service。当前前端回归基线为 29 项测试,TypeScript 类型检查和生产构建通过。
|
||||
|
||||
## 1. 结论
|
||||
|
||||
@@ -281,4 +281,4 @@ git diff --check passed
|
||||
|
||||
第三轮交互完善继续处理了文件切换、Markdown 选区格式、代码块默认状态、亮暗主题对比度和浮动工具栏失效问题。Provider 设置页增加 OpenAI、DeepSeek、Ollama 预设与模型自动发现,API Key 改为提交给后端加密保存,不进入 Pinia 或 Local Storage。智能体页面的运行状态、事件、工具、权限及导航文案已完成中文化,同时保留技术 ID 便于排障。
|
||||
|
||||
该轮新增 Store、Workspace、文件树、编辑器和中文标签回归测试;当前结果为前端 14 项、后端 71 项测试通过,生产构建通过。
|
||||
该轮新增 Store、Workspace、文件树、编辑器和中文标签回归测试;该轮当时结果为前端 14 项、后端 71 项测试通过,生产构建通过。最新全仓基线见本文开头的状态补充。
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
> 审阅范围:FastAPI、Knowledge / Retrieval Core、Agent Core、Extension Core、Provider Adapter、公共接口和后端开发文档。
|
||||
> 文档用途:记录问题形成原因、实际影响、修复判断和落地方案,供后续开发文档、比赛材料与技术博客使用。
|
||||
|
||||
> 2026-09-02 状态补充:本文记录的缺陷均保持修复。此后又加入 Provider 预设、模型发现、DeepSeek/OpenAI 凭据解析、Fernet 加密存储、Agent Trace 持久化、stdio MCP Plugin Host 和 Plugin Command/Settings,当前完整后端回归基线为 126 项测试通过。
|
||||
> 2026-09-02 状态补充:本文记录的缺陷均保持修复。此后又加入 Provider 预设、模型发现、DeepSeek/OpenAI 凭据解析、Fernet 加密存储、Agent Trace 持久化、stdio MCP Plugin Host 和 Plugin Command/Settings,当前完整后端回归基线为 136 项测试通过。
|
||||
|
||||
## 1. 审阅结论
|
||||
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
allowBuilds:
|
||||
esbuild: true
|
||||
@@ -0,0 +1,40 @@
|
||||
Lobe Icons — Copyright (c) 2023 LobeHub. MIT license; see LICENSE.
|
||||
Source: https://github.com/lobehub/lobe-icons
|
||||
Revision: 4aaf4ee1fb2678a7f989ea570f0f6ce14a9abf75
|
||||
Source directory: packages/static-svg/icons/
|
||||
Assets are bundled locally. Brand trademarks belong to their respective owners.
|
||||
File mapping (local: upstream):
|
||||
anthropic.svg: anthropic.svg
|
||||
baidu.svg: baidu-color.svg
|
||||
deepseek.svg: deepseek-color.svg
|
||||
hunyuan.svg: hunyuan-color.svg
|
||||
kimi.svg: kimi-color.svg
|
||||
minimax.svg: minimax-color.svg
|
||||
ollama.svg: ollama.svg
|
||||
openai.svg: openai.svg
|
||||
qwen.svg: qwen-color.svg
|
||||
siliconflow.svg: siliconcloud-color.svg
|
||||
stepfun.svg: stepfun-color.svg
|
||||
volcengine.svg: volcengine-color.svg
|
||||
zhipu.svg: zhipu-color.svg
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2023 LobeHub
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
@@ -0,0 +1,19 @@
|
||||
Lobe Icons — Copyright (c) 2023 LobeHub. MIT license; see LICENSE.
|
||||
Source: https://github.com/lobehub/lobe-icons
|
||||
Revision: 4aaf4ee1fb2678a7f989ea570f0f6ce14a9abf75
|
||||
Source directory: packages/static-svg/icons/
|
||||
Assets are bundled locally. Brand trademarks belong to their respective owners.
|
||||
File mapping (local: upstream):
|
||||
anthropic.svg: anthropic.svg
|
||||
baidu.svg: baidu-color.svg
|
||||
deepseek.svg: deepseek-color.svg
|
||||
hunyuan.svg: hunyuan-color.svg
|
||||
kimi.svg: kimi-color.svg
|
||||
minimax.svg: minimax-color.svg
|
||||
ollama.svg: ollama.svg
|
||||
openai.svg: openai.svg
|
||||
qwen.svg: qwen-color.svg
|
||||
siliconflow.svg: siliconcloud-color.svg
|
||||
stepfun.svg: stepfun-color.svg
|
||||
volcengine.svg: volcengine-color.svg
|
||||
zhipu.svg: zhipu-color.svg
|
||||
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2023 LobeHub
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
@@ -0,0 +1 @@
|
||||
<svg fill="currentColor" fill-rule="evenodd" height="1em" style="flex:none;line-height:1" viewBox="0 0 24 24" width="1em" xmlns="http://www.w3.org/2000/svg"><title>Anthropic</title><path d="M13.827 3.52h3.603L24 20h-3.603l-6.57-16.48zm-7.258 0h3.767L16.906 20h-3.674l-1.343-3.461H5.017l-1.344 3.46H0L6.57 3.522zm4.132 9.959L8.453 7.687 6.205 13.48H10.7z"></path></svg>
|
||||
|
After Width: | Height: | Size: 368 B |
@@ -0,0 +1 @@
|
||||
<svg height="1em" style="flex:none;line-height:1" viewBox="0 0 24 24" width="1em" xmlns="http://www.w3.org/2000/svg"><title>Baidu</title><path d="M8.859 11.735c1.017-1.71 4.059-3.083 6.202.286 1.579 2.284 4.284 4.397 4.284 4.397s2.027 1.601.73 4.684c-1.24 2.956-5.64 1.607-6.005 1.49l-.024-.009s-1.746-.568-3.776-.112c-2.026.458-3.773.286-3.773.286l-.045-.001c-.328-.01-2.38-.187-3.001-2.968-.675-3.028 2.365-4.687 2.592-4.968.226-.288 1.802-1.37 2.816-3.085zm.986 1.738v2.032h-1.64s-1.64.138-2.213 2.014c-.2 1.252.177 1.99.242 2.148.067.157.596 1.073 1.927 1.342h3.078v-7.514l-1.394-.022zm3.588 2.191l-1.44.024v3.956s.064.985 1.44 1.344h3.541v-5.3h-1.528v3.979h-1.46s-.466-.068-.553-.447v-3.556zM9.82 16.715v3.06H8.58s-.863-.045-1.126-1.049c-.136-.445.02-.959.088-1.16.063-.203.353-.671.951-.85H9.82zm9.525-9.036c2.086 0 2.646 2.06 2.646 2.742 0 .688.284 3.597-2.309 3.655-2.595.057-2.704-1.77-2.704-3.08 0-1.374.277-3.317 2.367-3.317zM4.24 6.08c1.523-.135 2.645 1.55 2.762 2.513.07.625.393 3.486-1.975 4-2.364.515-3.244-2.249-2.984-3.544 0 0 .28-2.797 2.197-2.969zm8.847-1.483c.14-1.31 1.69-3.316 2.931-3.028 1.236.285 2.367 1.944 2.137 3.37-.224 1.428-1.345 3.313-3.095 3.082-1.748-.226-2.143-1.823-1.973-3.424zM9.425 1c1.307 0 2.364 1.519 2.364 3.398 0 1.879-1.057 3.4-2.364 3.4s-2.367-1.521-2.367-3.4C7.058 2.518 8.118 1 9.425 1z" fill="#2932E1" fill-rule="nonzero"></path></svg>
|
||||
|
After Width: | Height: | Size: 1.4 KiB |
@@ -0,0 +1 @@
|
||||
<svg height="1em" style="flex:none;line-height:1" viewBox="0 0 24 24" width="1em" xmlns="http://www.w3.org/2000/svg"><title>DeepSeek</title><path d="M23.748 4.482c-.254-.124-.364.113-.512.234-.051.039-.094.09-.137.136-.372.397-.806.657-1.373.626-.829-.046-1.537.214-2.163.848-.133-.782-.575-1.248-1.247-1.548-.352-.156-.708-.311-.955-.65-.172-.241-.219-.51-.305-.774-.055-.16-.11-.323-.293-.35-.2-.031-.278.136-.356.276-.313.572-.434 1.202-.422 1.84.027 1.436.633 2.58 1.838 3.393.137.093.172.187.129.323-.082.28-.18.552-.266.833-.055.179-.137.217-.329.14a5.526 5.526 0 01-1.736-1.18c-.857-.828-1.631-1.742-2.597-2.458a11.365 11.365 0 00-.689-.471c-.985-.957.13-1.743.388-1.836.27-.098.093-.432-.779-.428-.872.004-1.67.295-2.687.684a3.055 3.055 0 01-.465.137 9.597 9.597 0 00-2.883-.102c-1.885.21-3.39 1.102-4.497 2.623C.082 8.606-.231 10.684.152 12.85c.403 2.284 1.569 4.175 3.36 5.653 1.858 1.533 3.997 2.284 6.438 2.14 1.482-.085 3.133-.284 4.994-1.86.47.234.962.327 1.78.397.63.059 1.236-.03 1.705-.128.735-.156.684-.837.419-.961-2.155-1.004-1.682-.595-2.113-.926 1.096-1.296 2.746-2.642 3.392-7.003.05-.347.007-.565 0-.845-.004-.17.035-.237.23-.256a4.173 4.173 0 001.545-.475c1.396-.763 1.96-2.015 2.093-3.517.02-.23-.004-.467-.247-.588zM11.581 18c-2.089-1.642-3.102-2.183-3.52-2.16-.392.024-.321.471-.235.763.09.288.207.486.371.739.114.167.192.416-.113.603-.673.416-1.842-.14-1.897-.167-1.361-.802-2.5-1.86-3.301-3.307-.774-1.393-1.224-2.887-1.298-4.482-.02-.386.093-.522.477-.592a4.696 4.696 0 011.529-.039c2.132.312 3.946 1.265 5.468 2.774.868.86 1.525 1.887 2.202 2.891.72 1.066 1.494 2.082 2.48 2.914.348.292.625.514.891.677-.802.09-2.14.11-3.054-.614zm1-6.44a.306.306 0 01.415-.287.302.302 0 01.2.288.306.306 0 01-.31.307.303.303 0 01-.304-.308zm3.11 1.596c-.2.081-.399.151-.59.16a1.245 1.245 0 01-.798-.254c-.274-.23-.47-.358-.552-.758a1.73 1.73 0 01.016-.588c.07-.327-.008-.537-.239-.727-.187-.156-.426-.199-.688-.199a.559.559 0 01-.254-.078c-.11-.054-.2-.19-.114-.358.028-.054.16-.186.192-.21.356-.202.767-.136 1.146.016.352.144.618.408 1.001.782.391.451.462.576.685.914.176.265.336.537.445.848.067.195-.019.354-.25.452z" fill="#4D6BFE"></path></svg>
|
||||
|
After Width: | Height: | Size: 2.1 KiB |
@@ -0,0 +1 @@
|
||||
<svg height="1em" style="flex:none;line-height:1" viewBox="0 0 24 24" width="1em" xmlns="http://www.w3.org/2000/svg"><title>Hunyuan</title><circle cx="12" cy="12" fill="#0055E9" r="12"></circle><path d="M12 0c.518 0 1.028.033 1.528.096A6.188 6.188 0 0112.12 12.28l-.12.001c-2.99 0-5.242 2.179-5.554 5.11-.223 2.086.353 4.412 2.242 6.146C3.672 22.1 0 17.479 0 12 0 5.373 5.373 0 12 0z" fill="#A8DFF5"></path><path d="M5.286 5a2.438 2.438 0 01.682 3.38c-3.962 5.966-3.215 10.743 2.648 15.136C3.636 22.056 0 17.452 0 12c0-1.787.39-3.482 1.09-5.006.253-.435.525-.872.817-1.311A2.438 2.438 0 015.286 5z" fill="#0055E9"></path><path d="M12.98.04c.272.021.543.053.81.093.583.106 1.117.254 1.538.44 6.638 2.927 8.07 10.052 1.748 15.642a4.125 4.125 0 01-5.822-.358c-1.51-1.706-1.3-4.184.357-5.822.858-.848 3.108-1.223 4.045-2.441 1.257-1.634 2.122-6.009-2.523-7.506L12.98.039z" fill="#00BCFF"></path><path d="M13.528.096A6.187 6.187 0 0112 12.281a5.75 5.75 0 00-1.71.255c.147-.905.595-1.784 1.321-2.501.858-.848 3.108-1.223 4.045-2.441 1.27-1.651 2.14-6.104-2.676-7.554.184.014.367.033.548.056z" fill="#ECECEE"></path></svg>
|
||||
|
After Width: | Height: | Size: 1.1 KiB |
@@ -0,0 +1 @@
|
||||
<svg height="1em" style="flex:none;line-height:1" viewBox="0 0 24 24" width="1em" xmlns="http://www.w3.org/2000/svg"><title>Kimi</title><path d="M21.846 0a1.923 1.923 0 110 3.846H20.15a.226.226 0 01-.227-.226V1.923C19.923.861 20.784 0 21.846 0z" fill="#1783FF"></path><path d="M11.065 11.199l7.257-7.2c.137-.136.06-.41-.116-.41H14.3a.164.164 0 00-.117.051l-7.82 7.756c-.122.12-.302.013-.302-.179V3.82c0-.127-.083-.23-.185-.23H3.186c-.103 0-.186.103-.186.23V19.77c0 .128.083.23.186.23h2.69c.103 0 .186-.102.186-.23v-3.25c0-.069.025-.135.069-.178l2.424-2.406a.158.158 0 01.205-.023l6.484 4.772a7.677 7.677 0 003.453 1.283c.108.012.2-.095.2-.23v-3.06c0-.117-.07-.212-.164-.227a5.028 5.028 0 01-2.027-.807l-5.613-4.064c-.117-.078-.132-.279-.028-.381z" fill="#fff"></path></svg>
|
||||
|
After Width: | Height: | Size: 773 B |
@@ -0,0 +1 @@
|
||||
<svg height="1em" style="flex:none;line-height:1" viewBox="0 0 24 24" width="1em" xmlns="http://www.w3.org/2000/svg"><title>Minimax</title><defs><linearGradient id="lobe-icons-minimax-_R_0_" x1="0%" x2="100.182%" y1="50.057%" y2="50.057%"><stop offset="0%" stop-color="#E2167E"></stop><stop offset="100%" stop-color="#FE603C"></stop></linearGradient></defs><path d="M16.278 2c1.156 0 2.093.927 2.093 2.07v12.501a.74.74 0 00.744.709.74.74 0 00.743-.709V9.099a2.06 2.06 0 012.071-2.049A2.06 2.06 0 0124 9.1v6.561a.649.649 0 01-.652.645.649.649 0 01-.653-.645V9.1a.762.762 0 00-.766-.758.762.762 0 00-.766.758v7.472a2.037 2.037 0 01-2.048 2.026 2.037 2.037 0 01-2.048-2.026v-12.5a.785.785 0 00-.788-.753.785.785 0 00-.789.752l-.001 15.904A2.037 2.037 0 0113.441 22a2.037 2.037 0 01-2.048-2.026V18.04c0-.356.292-.645.652-.645.36 0 .652.289.652.645v1.934c0 .263.142.506.372.638.23.131.514.131.744 0a.734.734 0 00.372-.638V4.07c0-1.143.937-2.07 2.093-2.07zm-5.674 0c1.156 0 2.093.927 2.093 2.07v11.523a.648.648 0 01-.652.645.648.648 0 01-.652-.645V4.07a.785.785 0 00-.789-.78.785.785 0 00-.789.78v14.013a2.06 2.06 0 01-2.07 2.048 2.06 2.06 0 01-2.071-2.048V9.1a.762.762 0 00-.766-.758.762.762 0 00-.766.758v3.8a2.06 2.06 0 01-2.071 2.049A2.06 2.06 0 010 12.9v-1.378c0-.357.292-.646.652-.646.36 0 .653.29.653.646V12.9c0 .418.343.757.766.757s.766-.339.766-.757V9.099a2.06 2.06 0 012.07-2.048 2.06 2.06 0 012.071 2.048v8.984c0 .419.343.758.767.758.423 0 .766-.339.766-.758V4.07c0-1.143.937-2.07 2.093-2.07z" fill="url(#lobe-icons-minimax-_R_0_)" fill-rule="nonzero"></path></svg>
|
||||
|
After Width: | Height: | Size: 1.5 KiB |
@@ -0,0 +1 @@
|
||||
<svg fill="currentColor" fill-rule="evenodd" height="1em" style="flex:none;line-height:1" viewBox="0 0 24 24" width="1em" xmlns="http://www.w3.org/2000/svg"><title>Ollama</title><path d="M7.905 1.09c.216.085.411.225.588.41.295.306.544.744.734 1.263.191.522.315 1.1.362 1.68a5.054 5.054 0 012.049-.636l.051-.004c.87-.07 1.73.087 2.48.474.101.053.2.11.297.17.05-.569.172-1.134.36-1.644.19-.52.439-.957.733-1.264a1.67 1.67 0 01.589-.41c.257-.1.53-.118.796-.042.401.114.745.368 1.016.737.248.337.434.769.561 1.287.23.934.27 2.163.115 3.645l.053.04.026.019c.757.576 1.284 1.397 1.563 2.35.435 1.487.216 3.155-.534 4.088l-.018.021.002.003c.417.762.67 1.567.724 2.4l.002.03c.064 1.065-.2 2.137-.814 3.19l-.007.01.01.024c.472 1.157.62 2.322.438 3.486l-.006.039a.651.651 0 01-.747.536.648.648 0 01-.54-.742c.167-1.033.01-2.069-.48-3.123a.643.643 0 01.04-.617l.004-.006c.604-.924.854-1.83.8-2.72-.046-.779-.325-1.544-.8-2.273a.644.644 0 01.18-.886l.009-.006c.243-.159.467-.565.58-1.12a4.229 4.229 0 00-.095-1.974c-.205-.7-.58-1.284-1.105-1.683-.595-.454-1.383-.673-2.38-.61a.653.653 0 01-.632-.371c-.314-.665-.772-1.141-1.343-1.436a3.288 3.288 0 00-1.772-.332c-1.245.099-2.343.801-2.67 1.686a.652.652 0 01-.61.425c-1.067.002-1.893.252-2.497.703-.522.39-.878.935-1.066 1.588a4.07 4.07 0 00-.068 1.886c.112.558.331 1.02.582 1.269l.008.007c.212.207.257.53.109.785-.36.622-.629 1.549-.673 2.44-.05 1.018.186 1.902.719 2.536l.016.019a.643.643 0 01.095.69c-.576 1.236-.753 2.252-.562 3.052a.652.652 0 01-1.269.298c-.243-1.018-.078-2.184.473-3.498l.014-.035-.008-.012a4.339 4.339 0 01-.598-1.309l-.005-.019a5.764 5.764 0 01-.177-1.785c.044-.91.278-1.842.622-2.59l.012-.026-.002-.002c-.293-.418-.51-.953-.63-1.545l-.005-.024a5.352 5.352 0 01.093-2.49c.262-.915.777-1.701 1.536-2.269.06-.045.123-.09.186-.132-.159-1.493-.119-2.73.112-3.67.127-.518.314-.95.562-1.287.27-.368.614-.622 1.015-.737.266-.076.54-.059.797.042zm4.116 9.09c.936 0 1.8.313 2.446.855.63.527 1.005 1.235 1.005 1.94 0 .888-.406 1.58-1.133 2.022-.62.375-1.451.557-2.403.557-1.009 0-1.871-.259-2.493-.734-.617-.47-.963-1.13-.963-1.845 0-.707.398-1.417 1.056-1.946.668-.537 1.55-.849 2.485-.849zm0 .896a3.07 3.07 0 00-1.916.65c-.461.37-.722.835-.722 1.25 0 .428.21.829.61 1.134.455.347 1.124.548 1.943.548.799 0 1.473-.147 1.932-.426.463-.28.7-.686.7-1.257 0-.423-.246-.89-.683-1.256-.484-.405-1.14-.643-1.864-.643zm.662 1.21l.004.004c.12.151.095.37-.056.49l-.292.23v.446a.375.375 0 01-.376.373.375.375 0 01-.376-.373v-.46l-.271-.218a.347.347 0 01-.052-.49.353.353 0 01.494-.051l.215.172.22-.174a.353.353 0 01.49.051zm-5.04-1.919c.478 0 .867.39.867.871a.87.87 0 01-.868.871.87.87 0 01-.867-.87.87.87 0 01.867-.872zm8.706 0c.48 0 .868.39.868.871a.87.87 0 01-.868.871.87.87 0 01-.867-.87.87.87 0 01.867-.872zM7.44 2.3l-.003.002a.659.659 0 00-.285.238l-.005.006c-.138.189-.258.467-.348.832-.17.692-.216 1.631-.124 2.782.43-.128.899-.208 1.404-.237l.01-.001.019-.034c.046-.082.095-.161.148-.239.123-.771.022-1.692-.253-2.444-.134-.364-.297-.65-.453-.813a.628.628 0 00-.107-.09L7.44 2.3zm9.174.04l-.002.001a.628.628 0 00-.107.09c-.156.163-.32.45-.453.814-.29.794-.387 1.776-.23 2.572l.058.097.008.014h.03a5.184 5.184 0 011.466.212c.086-1.124.038-2.043-.128-2.722-.09-.365-.21-.643-.349-.832l-.004-.006a.659.659 0 00-.285-.239h-.004z"></path></svg>
|
||||
|
After Width: | Height: | Size: 3.2 KiB |
@@ -0,0 +1 @@
|
||||
<svg fill="currentColor" fill-rule="evenodd" height="1em" style="flex:none;line-height:1" viewBox="0 0 24 24" width="1em" xmlns="http://www.w3.org/2000/svg"><title>OpenAI</title><path d="M9.205 8.658v-2.26c0-.19.072-.333.238-.428l4.543-2.616c.619-.357 1.356-.523 2.117-.523 2.854 0 4.662 2.212 4.662 4.566 0 .167 0 .357-.024.547l-4.71-2.759a.797.797 0 00-.856 0l-5.97 3.473zm10.609 8.8V12.06c0-.333-.143-.57-.429-.737l-5.97-3.473 1.95-1.118a.433.433 0 01.476 0l4.543 2.617c1.309.76 2.189 2.378 2.189 3.948 0 1.808-1.07 3.473-2.76 4.163zM7.802 12.703l-1.95-1.142c-.167-.095-.239-.238-.239-.428V5.899c0-2.545 1.95-4.472 4.591-4.472 1 0 1.927.333 2.712.928L8.23 5.067c-.285.166-.428.404-.428.737v6.898zM12 15.128l-2.795-1.57v-3.33L12 8.658l2.795 1.57v3.33L12 15.128zm1.796 7.23c-1 0-1.927-.332-2.712-.927l4.686-2.712c.285-.166.428-.404.428-.737v-6.898l1.974 1.142c.167.095.238.238.238.428v5.233c0 2.545-1.974 4.472-4.614 4.472zm-5.637-5.303l-4.544-2.617c-1.308-.761-2.188-2.378-2.188-3.948A4.482 4.482 0 014.21 6.327v5.423c0 .333.143.571.428.738l5.947 3.449-1.95 1.118a.432.432 0 01-.476 0zm-.262 3.9c-2.688 0-4.662-2.021-4.662-4.519 0-.19.024-.38.047-.57l4.686 2.71c.286.167.571.167.856 0l5.97-3.448v2.26c0 .19-.07.333-.237.428l-4.543 2.616c-.619.357-1.356.523-2.117.523zm5.899 2.83a5.947 5.947 0 005.827-4.756C22.287 18.339 24 15.84 24 13.296c0-1.665-.713-3.282-1.998-4.448.119-.5.19-.999.19-1.498 0-3.401-2.759-5.947-5.946-5.947-.642 0-1.26.095-1.88.31A5.962 5.962 0 0010.205 0a5.947 5.947 0 00-5.827 4.757C1.713 5.447 0 7.945 0 10.49c0 1.666.713 3.283 1.998 4.448-.119.5-.19 1-.19 1.499 0 3.401 2.759 5.946 5.946 5.946.642 0 1.26-.095 1.88-.309a5.96 5.96 0 004.162 1.713z"></path></svg>
|
||||
|
After Width: | Height: | Size: 1.6 KiB |
@@ -0,0 +1 @@
|
||||
<svg height="1em" style="flex:none;line-height:1" viewBox="0 0 24 24" width="1em" xmlns="http://www.w3.org/2000/svg"><title>Qwen</title><path d="M12.604 1.34c.393.69.784 1.382 1.174 2.075a.18.18 0 00.157.091h5.552c.174 0 .322.11.446.327l1.454 2.57c.19.337.24.478.024.837-.26.43-.513.864-.76 1.3l-.367.658c-.106.196-.223.28-.04.512l2.652 4.637c.172.301.111.494-.043.77-.437.785-.882 1.564-1.335 2.34-.159.272-.352.375-.68.37-.777-.016-1.552-.01-2.327.016a.099.099 0 00-.081.05 575.097 575.097 0 01-2.705 4.74c-.169.293-.38.363-.725.364-.997.003-2.002.004-3.017.002a.537.537 0 01-.465-.271l-1.335-2.323a.09.09 0 00-.083-.049H4.982c-.285.03-.553-.001-.805-.092l-1.603-2.77a.543.543 0 01-.002-.54l1.207-2.12a.198.198 0 000-.197 550.951 550.951 0 01-1.875-3.272l-.79-1.395c-.16-.31-.173-.496.095-.965.465-.813.927-1.625 1.387-2.436.132-.234.304-.334.584-.335a338.3 338.3 0 012.589-.001.124.124 0 00.107-.063l2.806-4.895a.488.488 0 01.422-.246c.524-.001 1.053 0 1.583-.006L11.704 1c.341-.003.724.032.9.34zm-3.432.403a.06.06 0 00-.052.03L6.254 6.788a.157.157 0 01-.135.078H3.253c-.056 0-.07.025-.041.074l5.81 10.156c.025.042.013.062-.034.063l-2.795.015a.218.218 0 00-.2.116l-1.32 2.31c-.044.078-.021.118.068.118l5.716.008c.046 0 .08.02.104.061l1.403 2.454c.046.081.092.082.139 0l5.006-8.76.783-1.382a.055.055 0 01.096 0l1.424 2.53a.122.122 0 00.107.062l2.763-.02a.04.04 0 00.035-.02.041.041 0 000-.04l-2.9-5.086a.108.108 0 010-.113l.293-.507 1.12-1.977c.024-.041.012-.062-.035-.062H9.2c-.059 0-.073-.026-.043-.077l1.434-2.505a.107.107 0 000-.114L9.225 1.774a.06.06 0 00-.053-.031zm6.29 8.02c.046 0 .058.02.034.06l-.832 1.465-2.613 4.585a.056.056 0 01-.05.029.058.058 0 01-.05-.029L8.498 9.841c-.02-.034-.01-.052.028-.054l.216-.012 6.722-.012z" fill="url(#lobe-icons-qwen-_R_0_)" fill-rule="nonzero"></path><defs><linearGradient id="lobe-icons-qwen-_R_0_" x1="0%" x2="100%" y1="0%" y2="0%"><stop offset="0%" stop-color="#6336E7" stop-opacity=".84"></stop><stop offset="100%" stop-color="#6F69F7" stop-opacity=".84"></stop></linearGradient></defs></svg>
|
||||
|
After Width: | Height: | Size: 2.0 KiB |
@@ -0,0 +1 @@
|
||||
<svg height="1em" style="flex:none;line-height:1" viewBox="0 0 24 24" width="1em" xmlns="http://www.w3.org/2000/svg"><title>SiliconCloud</title><path clip-rule="evenodd" d="M22.956 6.521H12.522c-.577 0-1.044.468-1.044 1.044v3.13c0 .577-.466 1.044-1.043 1.044H1.044c-.577 0-1.044.467-1.044 1.044v4.174C0 17.533.467 18 1.044 18h10.434c.577 0 1.044-.467 1.044-1.043v-3.13c0-.578.466-1.044 1.043-1.044h9.391c.577 0 1.044-.467 1.044-1.044V7.565c0-.576-.467-1.044-1.044-1.044z" fill="#6E29F6" fill-rule="evenodd"></path></svg>
|
||||
|
After Width: | Height: | Size: 520 B |
@@ -0,0 +1 @@
|
||||
<svg height="1em" style="flex:none;line-height:1" viewBox="0 0 24 24" width="1em" xmlns="http://www.w3.org/2000/svg"><title>Stepfun</title><path d="M22.012 0h1.032v.927H24v.968h-.956V3.78h-1.032V1.896h-1.878v-.97h1.878V0zM2.6 12.371V1.87h.969v10.502h-.97zm10.423.66h10.95v.918h-6.208v9.579h-4.742V13.03zM5.629 3.333v12.356H0v4.51h10.386V8L20.859 8l-.003-4.668-15.227.001z" fill="url(#lobe-icons-stepfun-_R_0_)" fill-rule="evenodd"></path><defs><linearGradient gradientUnits="userSpaceOnUse" id="lobe-icons-stepfun-_R_0_" x1="1.646" x2="18.342" y1="1.916" y2="22.091"><stop stop-color="#01A9FF"></stop><stop offset="1" stop-color="#0160FF"></stop></linearGradient></defs></svg>
|
||||
|
After Width: | Height: | Size: 676 B |
@@ -0,0 +1 @@
|
||||
<svg height="1em" style="flex:none;line-height:1" viewBox="0 0 24 24" width="1em" xmlns="http://www.w3.org/2000/svg"><title>Volcengine</title><path d="M19.44 10.153l-2.936 11.586a.215.215 0 00.214.261h5.87a.215.215 0 00.214-.261l-2.95-11.586a.214.214 0 00-.412 0zM3.28 12.778l-2.275 8.96A.214.214 0 001.22 22h4.532a.212.212 0 00.214-.165.214.214 0 000-.097l-2.276-8.96a.214.214 0 00-.41 0z" fill="#00E5E5"></path><path d="M7.29 5.359L3.148 21.738a.215.215 0 00.203.261h8.29a.214.214 0 00.215-.261L7.7 5.358a.214.214 0 00-.41 0z" fill="#006EFF"></path><path d="M14.44.15a.214.214 0 00-.41 0L8.366 21.739a.214.214 0 00.214.261H19.9a.216.216 0 00.171-.078.214.214 0 00.044-.183L14.439.15z" fill="#006EFF"></path><path d="M10.278 7.741L6.685 21.736a.214.214 0 00.214.264h7.17a.215.215 0 00.214-.264L10.688 7.741a.214.214 0 00-.41 0z" fill="#00E5E5"></path></svg>
|
||||
|
After Width: | Height: | Size: 858 B |
@@ -0,0 +1 @@
|
||||
<svg height="1em" style="flex:none;line-height:1" viewBox="0 0 24 24" width="1em" xmlns="http://www.w3.org/2000/svg"><title>Zhipu</title><path d="M11.991 23.503a.24.24 0 00-.244.248.24.24 0 00.244.249.24.24 0 00.245-.249.24.24 0 00-.22-.247l-.025-.001zM9.671 5.365a1.697 1.697 0 011.099 2.132l-.071.172-.016.04-.018.054c-.07.16-.104.32-.104.498-.035.71.47 1.279 1.186 1.314h.366c1.309.053 2.338 1.173 2.286 2.523-.052 1.332-1.152 2.38-2.478 2.327h-.174c-.715.018-1.274.64-1.239 1.368 0 .124.018.23.053.337.209.373.54.658.96.8.75.23 1.517-.125 1.9-.782l.018-.035c.402-.64 1.17-.96 1.92-.711.854.284 1.378 1.226 1.099 2.167a1.661 1.661 0 01-2.077 1.102 1.711 1.711 0 01-.907-.711l-.017-.035c-.2-.323-.463-.58-.851-.711l-.056-.018a1.646 1.646 0 00-1.954.746 1.66 1.66 0 01-1.065.764 1.677 1.677 0 01-1.989-1.279c-.209-.906.332-1.83 1.257-2.043a1.51 1.51 0 01.296-.035h.018c.68-.071 1.151-.622 1.116-1.333a1.307 1.307 0 00-.227-.693 2.515 2.515 0 01-.366-1.403 2.39 2.39 0 01.366-1.208c.14-.195.21-.444.227-.693.018-.71-.506-1.261-1.186-1.332l-.07-.018a1.43 1.43 0 01-.299-.07l-.05-.019a1.7 1.7 0 01-1.047-2.114 1.68 1.68 0 012.094-1.101zm-5.575 10.11c.26-.264.639-.367.994-.27.355.096.633.379.728.74.095.362-.007.748-.267 1.013-.402.41-1.053.41-1.455 0a1.062 1.062 0 010-1.482zm14.845-.294c.359-.09.738.024.992.297.254.274.344.665.237 1.025-.107.36-.396.634-.756.718-.551.128-1.1-.22-1.23-.781a1.05 1.05 0 01.757-1.26zm-.064-4.39c.314.32.49.753.49 1.206 0 .452-.176.886-.49 1.206-.315.32-.74.5-1.185.5-.444 0-.87-.18-1.184-.5a1.727 1.727 0 010-2.412 1.654 1.654 0 012.369 0zm-11.243.163c.364.484.447 1.128.218 1.691a1.665 1.665 0 01-2.188.923c-.855-.36-1.26-1.358-.907-2.228a1.68 1.68 0 011.33-1.038c.593-.08 1.183.169 1.547.652zm11.545-4.221c.368 0 .708.2.892.524.184.324.184.724 0 1.048a1.026 1.026 0 01-.892.524c-.568 0-1.03-.47-1.03-1.048 0-.579.462-1.048 1.03-1.048zm-14.358 0c.368 0 .707.2.891.524.184.324.184.724 0 1.048a1.026 1.026 0 01-.891.524c-.569 0-1.03-.47-1.03-1.048 0-.579.461-1.048 1.03-1.048zm10.031-1.475c.925 0 1.675.764 1.675 1.706s-.75 1.705-1.675 1.705-1.674-.763-1.674-1.705c0-.942.75-1.706 1.674-1.706zm-2.626-.684c.362-.082.653-.356.761-.718a1.062 1.062 0 00-.238-1.028 1.017 1.017 0 00-.996-.294c-.547.14-.881.7-.752 1.257.13.558.675.907 1.225.783zm0 16.876c.359-.087.644-.36.75-.72a1.062 1.062 0 00-.237-1.019 1.018 1.018 0 00-.985-.301 1.037 1.037 0 00-.762.717c-.108.361-.017.754.239 1.028.245.263.606.377.953.305l.043-.01zM17.19 3.5a.631.631 0 00.628-.64c0-.355-.279-.64-.628-.64a.631.631 0 00-.628.64c0 .355.28.64.628.64zm-10.38 0a.631.631 0 00.628-.64c0-.355-.28-.64-.628-.64a.631.631 0 00-.628.64c0 .355.279.64.628.64zm-5.182 7.852a.631.631 0 00-.628.64c0 .354.28.639.628.639a.63.63 0 00.627-.606l.001-.034a.62.62 0 00-.628-.64zm5.182 9.13a.631.631 0 00-.628.64c0 .355.279.64.628.64a.631.631 0 00.628-.64c0-.355-.28-.64-.628-.64zm10.38.018a.631.631 0 00-.628.64c0 .355.28.64.628.64a.631.631 0 00.628-.64c0-.355-.279-.64-.628-.64zm5.182-9.148a.631.631 0 00-.628.64c0 .354.279.639.628.639a.631.631 0 00.628-.64c0-.355-.28-.64-.628-.64zm-.384-4.992a.24.24 0 00.244-.249.24.24 0 00-.244-.249.24.24 0 00-.244.249c0 .142.122.249.244.249zM11.991.497a.24.24 0 00.245-.248A.24.24 0 0011.99 0a.24.24 0 00-.244.249c0 .133.108.236.223.247l.021.001zM2.011 6.36a.24.24 0 00.245-.249.24.24 0 00-.244-.249.24.24 0 00-.244.249.24.24 0 00.244.249zm0 11.263a.24.24 0 00-.243.248.24.24 0 00.244.249.24.24 0 00.244-.249.252.252 0 00-.244-.248zm19.995-.018a.24.24 0 00-.245.248.24.24 0 00.245.25.24.24 0 00.244-.25.252.252 0 00-.244-.248z" fill="#3859FF" fill-rule="nonzero"></path></svg>
|
||||
|
After Width: | Height: | Size: 3.5 KiB |
@@ -0,0 +1,77 @@
|
||||
// @vitest-environment happy-dom
|
||||
import { flushPromises, mount } from '@vue/test-utils'
|
||||
import { createPinia, setActivePinia } from 'pinia'
|
||||
import { createMemoryHistory, createRouter } from 'vue-router'
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import * as pluginService from '@/services/pluginService'
|
||||
import { useEditorStore } from '@/stores/editor'
|
||||
import { useWorkspaceStore } from '@/stores/workspace'
|
||||
import CommandPalette from './CommandPalette.vue'
|
||||
|
||||
vi.mock('@/services/pluginService', async (loadOriginal) => {
|
||||
const original = await loadOriginal<typeof import('@/services/pluginService')>()
|
||||
return { ...original, listPluginCommands: vi.fn(), executePluginCommand: vi.fn() }
|
||||
})
|
||||
|
||||
beforeEach(() => {
|
||||
setActivePinia(createPinia())
|
||||
vi.mocked(pluginService.listPluginCommands).mockResolvedValue([{
|
||||
command_id: 'demo.selection',
|
||||
plugin_id: 'demo',
|
||||
title: '处理选区',
|
||||
description: '',
|
||||
icon: null,
|
||||
locations: ['command_palette'],
|
||||
when: ['workspace.has_vault', 'editor.has_note', 'editor.has_selection'],
|
||||
parameters: { type: 'object', properties: {}, additionalProperties: false },
|
||||
enabled: true,
|
||||
}])
|
||||
vi.mocked(pluginService.executePluginCommand).mockResolvedValue({
|
||||
command_id: 'demo.selection',
|
||||
status: 'completed',
|
||||
effect: { type: 'notification', payload: { level: 'success', message: '完成' } },
|
||||
})
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
document.body.innerHTML = ''
|
||||
vi.restoreAllMocks()
|
||||
})
|
||||
|
||||
describe('CommandPalette Plugin Command', () => {
|
||||
it('filters by when context and sends stable backend identities plus the captured selection', async () => {
|
||||
const workspace = useWorkspaceStore()
|
||||
workspace.hasVault = true
|
||||
workspace.vaultId = 'vault-default'
|
||||
const editor = useEditorStore()
|
||||
editor.currentNoteId = 'note-1'
|
||||
editor.currentFilePath = '/note.md'
|
||||
|
||||
vi.spyOn(window, 'getSelection').mockReturnValue({
|
||||
toString: () => 'selected text',
|
||||
} as Selection)
|
||||
|
||||
const router = createRouter({
|
||||
history: createMemoryHistory(),
|
||||
routes: [{ path: '/', component: { template: '<div />' } }],
|
||||
})
|
||||
await router.push('/')
|
||||
const wrapper = mount(CommandPalette, { attachTo: document.body, global: { plugins: [router] } })
|
||||
|
||||
window.dispatchEvent(new KeyboardEvent('keydown', { key: 'p', ctrlKey: true }))
|
||||
await flushPromises()
|
||||
const command = Array.from(document.querySelectorAll('button')).find((button) => button.textContent?.includes('处理选区'))
|
||||
expect(command).toBeTruthy()
|
||||
command!.click()
|
||||
await flushPromises()
|
||||
|
||||
expect(pluginService.executePluginCommand).toHaveBeenCalledWith('demo.selection', {}, {
|
||||
vault_id: 'vault-default',
|
||||
note_id: 'note-1',
|
||||
file_path: '/note.md',
|
||||
selection: 'selected text',
|
||||
})
|
||||
expect(document.body.textContent).toContain('完成')
|
||||
wrapper.unmount()
|
||||
})
|
||||
})
|
||||
@@ -5,18 +5,26 @@ import { useEditorStore } from '@/stores/editor'
|
||||
import { useThemeStore } from '@/stores/theme'
|
||||
import { useWorkspaceStore } from '@/stores/workspace'
|
||||
import * as workspaceService from '@/services/workspaceService'
|
||||
import * as pluginService from '@/services/pluginService'
|
||||
import type { PluginCommand, PluginCommandEffect } from '@/contracts'
|
||||
import { usePluginStore } from '@/stores/plugin'
|
||||
|
||||
const router = useRouter()
|
||||
const editorStore = useEditorStore()
|
||||
const themeStore = useThemeStore()
|
||||
const workspaceStore = useWorkspaceStore()
|
||||
const pluginStore = usePluginStore()
|
||||
const open = ref(false)
|
||||
const query = ref('')
|
||||
const input = ref<HTMLInputElement | null>(null)
|
||||
const pluginCommands = ref<PluginCommand[]>([])
|
||||
const commandError = ref('')
|
||||
const commandNotice = ref('')
|
||||
const selectionSnapshot = ref<string | null>(null)
|
||||
|
||||
interface Command { id: string; label: string; hint: string; run: () => void | Promise<void> }
|
||||
|
||||
const commands = computed<Command[]>(() => [
|
||||
const builtinCommands = computed<Command[]>(() => [
|
||||
{ id: 'workspace', label: '打开工作区', hint: '导航', run: () => router.push('/workspace') },
|
||||
{ id: 'search', label: '全局搜索', hint: '导航', run: () => router.push('/search') },
|
||||
{ id: 'chat', label: '打开 AI 对话', hint: '导航', run: () => router.push('/chat') },
|
||||
@@ -28,14 +36,37 @@ const commands = computed<Command[]>(() => [
|
||||
{ id: 'new-note', label: '创建笔记', hint: '工作区', run: createNote },
|
||||
])
|
||||
|
||||
const commands = computed<Command[]>(() => [
|
||||
...builtinCommands.value,
|
||||
...pluginCommands.value.filter(isPluginCommandAvailable).map((command) => ({
|
||||
id: 'plugin:' + command.command_id,
|
||||
label: command.title,
|
||||
hint: 'Plugin · ' + command.plugin_id,
|
||||
run: () => executePluginCommand(command),
|
||||
})),
|
||||
])
|
||||
|
||||
function isPluginCommandAvailable(command: PluginCommand) {
|
||||
if (!command.enabled) return false
|
||||
return command.when.every((condition) => {
|
||||
if (condition === 'workspace.has_vault') return Boolean(workspaceStore.vaultId)
|
||||
if (condition === 'editor.has_note') return Boolean(editorStore.currentNoteId)
|
||||
if (condition === 'editor.has_selection') return Boolean(selectionSnapshot.value)
|
||||
return false
|
||||
})
|
||||
}
|
||||
|
||||
const filteredCommands = computed(() => {
|
||||
const value = query.value.trim().toLocaleLowerCase()
|
||||
return value ? commands.value.filter((command) => `${command.label} ${command.hint}`.toLocaleLowerCase().includes(value)) : commands.value
|
||||
})
|
||||
|
||||
function show() {
|
||||
selectionSnapshot.value = window.getSelection()?.toString() || null
|
||||
open.value = true
|
||||
query.value = ''
|
||||
commandError.value = ''
|
||||
void loadPluginCommands()
|
||||
void nextTick(() => input.value?.focus())
|
||||
}
|
||||
|
||||
@@ -44,7 +75,11 @@ function hide() { open.value = false }
|
||||
async function execute(command: Command | undefined) {
|
||||
if (!command) return
|
||||
hide()
|
||||
try {
|
||||
await command.run()
|
||||
} catch (error) {
|
||||
commandNotice.value = error instanceof Error ? error.message : '命令执行失败'
|
||||
}
|
||||
}
|
||||
|
||||
async function createNote() {
|
||||
@@ -58,6 +93,55 @@ async function createNote() {
|
||||
await router.push('/workspace')
|
||||
}
|
||||
|
||||
async function loadPluginCommands() {
|
||||
try {
|
||||
pluginCommands.value = await pluginService.listPluginCommands('command_palette')
|
||||
} catch (error) {
|
||||
commandError.value = error instanceof Error ? error.message : 'Plugin 命令加载失败'
|
||||
}
|
||||
}
|
||||
|
||||
function hasRequiredArguments(command: PluginCommand) {
|
||||
return Array.isArray(command.parameters.required) && command.parameters.required.length > 0
|
||||
}
|
||||
|
||||
async function executePluginCommand(command: PluginCommand) {
|
||||
if (hasRequiredArguments(command)) {
|
||||
pluginStore.selectPlugin(command.plugin_id)
|
||||
await router.push('/extensions/plugins')
|
||||
commandNotice.value = '请在 Plugin 详情页填写参数后执行“' + command.title + '”。'
|
||||
return
|
||||
}
|
||||
const result = await pluginService.executePluginCommand(command.command_id, {}, {
|
||||
vault_id: workspaceStore.hasVault ? workspaceStore.vaultId : null,
|
||||
note_id: editorStore.currentNoteId,
|
||||
file_path: editorStore.currentFilePath,
|
||||
selection: selectionSnapshot.value,
|
||||
})
|
||||
await applyPluginEffect(result.effect)
|
||||
}
|
||||
|
||||
async function applyPluginEffect(effect: PluginCommandEffect) {
|
||||
if (effect.type === 'notification') { commandNotice.value = effect.payload.message; return }
|
||||
if (effect.type === 'navigate') {
|
||||
const routes: Record<string, string> = {
|
||||
'vault-entry': '/', workspace: '/workspace', search: '/search', chat: '/chat',
|
||||
agent: '/agent/runs', tasks: '/tasks', skills: '/extensions/skills',
|
||||
plugins: '/extensions/plugins', themes: '/themes', settings: '/settings',
|
||||
}
|
||||
await router.push(routes[effect.payload.route])
|
||||
return
|
||||
}
|
||||
if (effect.type === 'refresh') {
|
||||
if (effect.payload.scope === 'plugins') await pluginStore.loadPlugins()
|
||||
if (effect.payload.scope === 'commands') await loadPluginCommands()
|
||||
commandNotice.value = '相关数据已刷新。'
|
||||
return
|
||||
}
|
||||
if (effect.type === 'job') { commandNotice.value = '后台任务已创建:' + effect.payload.job_id; return }
|
||||
commandNotice.value = 'Plugin 命令执行完成。'
|
||||
}
|
||||
|
||||
function handleKeydown(event: KeyboardEvent) {
|
||||
if ((event.ctrlKey || event.metaKey) && event.key.toLocaleLowerCase() === 'p') {
|
||||
event.preventDefault()
|
||||
@@ -72,10 +156,14 @@ onBeforeUnmount(() => window.removeEventListener('keydown', handleKeydown))
|
||||
</script>
|
||||
|
||||
<template>
|
||||
<div v-if="commandNotice" class="command-toast" role="status">
|
||||
<span>{{ commandNotice }}</span><button aria-label="关闭通知" @click="commandNotice = ''">×</button>
|
||||
</div>
|
||||
<Teleport to="body">
|
||||
<div v-if="open" class="command-backdrop" @click.self="hide">
|
||||
<section class="command-palette" role="dialog" aria-modal="true" aria-label="命令面板">
|
||||
<input ref="input" v-model="query" class="command-input" placeholder="输入命令…" @keydown.enter.prevent="execute(filteredCommands[0])" />
|
||||
<p v-if="commandError" class="command-error">{{ commandError }}</p>
|
||||
<div class="command-list">
|
||||
<button v-for="command in filteredCommands" :key="command.id" type="button" @click="execute(command)">
|
||||
<span>{{ command.label }}</span><small>{{ command.hint }}</small>
|
||||
@@ -99,6 +187,9 @@ onBeforeUnmount(() => window.removeEventListener('keydown', handleKeydown))
|
||||
.command-list small, .command-list p, footer { color: var(--color-text-tertiary); }
|
||||
.command-list p { padding: var(--space-xl); text-align: center; }
|
||||
footer { display: flex; gap: var(--space-lg); padding: var(--space-sm) var(--space-lg); border-top: 1px solid var(--color-border-subtle); font-size: var(--font-size-xs); }
|
||||
.command-error { margin: var(--space-sm); padding: var(--space-sm) var(--space-md); border-radius: var(--radius-md); background: var(--color-error-soft); color: var(--color-error); font-size: var(--font-size-sm); }
|
||||
.command-toast { position: fixed; top: 48px; right: var(--space-xl); z-index: calc(var(--z-modal) + 1); display: flex; align-items: center; gap: var(--space-lg); max-width: min(420px, calc(100vw - 32px)); padding: var(--space-md) var(--space-lg); border: 1px solid var(--color-border-default); border-radius: var(--radius-lg); background: var(--color-surface-elevated); box-shadow: var(--shadow-lg); animation: notice-in var(--motion-normal) both; }
|
||||
.command-toast button { color: var(--color-text-tertiary); font-size: var(--font-size-xl); }
|
||||
|
||||
@keyframes command-backdrop-in { from { opacity: 0; } to { opacity: 1; } }
|
||||
@keyframes command-palette-in { from { opacity: 0; transform: translateY(-8px) scale(.99); } to { opacity: 1; transform: translateY(0) scale(1); } }
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
<script setup lang="ts">
|
||||
import { useRoute, useRouter } from 'vue-router'
|
||||
import { computed, ref } from 'vue'
|
||||
import { ArrowLeftBold, ArrowRightBold, Brush, ChatDotRound, CircleCheck, Connection, Cpu, FolderOpened, Lightning, Search, Setting } from '@element-plus/icons-vue'
|
||||
import { ArrowLeftBold, ArrowRightBold, Brush, ChatDotRound, CircleCheck, Connection, Cpu, FolderOpened, Lightning, Monitor, Search, Setting } from '@element-plus/icons-vue'
|
||||
import AppIcon from './AppIcon.vue'
|
||||
|
||||
const route = useRoute()
|
||||
@@ -16,6 +16,7 @@ const navItems = [
|
||||
{ name: 'tasks', icon: CircleCheck, label: '任务' },
|
||||
{ name: 'skills', icon: Lightning, label: 'Skill' },
|
||||
{ name: 'plugins', icon: Connection, label: 'Plugin' },
|
||||
{ name: 'mcp-servers', icon: Monitor, label: 'MCP' },
|
||||
{ name: 'themes', icon: Brush, label: '主题' },
|
||||
{ name: 'settings', icon: Setting, label: '设置' },
|
||||
]
|
||||
|
||||
@@ -21,7 +21,7 @@ const pageTitle = computed(() => {
|
||||
agent: '智能体执行轨迹',
|
||||
tasks: '任务',
|
||||
skills: 'Skill 管理',
|
||||
plugins: 'Plugin 管理',
|
||||
plugins: 'Plugin 与 MCP',
|
||||
themes: '主题管理',
|
||||
settings: '设置',
|
||||
}
|
||||
|
||||
@@ -193,7 +193,7 @@ export interface ToolDefinition {
|
||||
name: string
|
||||
description: string
|
||||
parameters: Record<string, unknown>
|
||||
source?: 'builtin' | 'plugin'
|
||||
source?: 'builtin' | 'plugin' | 'mcp_server'
|
||||
plugin_id?: string
|
||||
}
|
||||
|
||||
@@ -416,6 +416,39 @@ export interface ProviderPreset {
|
||||
base_url: string
|
||||
default_credential_id?: string | null
|
||||
requires_credential: boolean
|
||||
logo_id?: string
|
||||
description?: string
|
||||
capabilities?: string[]
|
||||
}
|
||||
|
||||
export type ProviderUpdateRequest = Partial<Omit<ProviderConfig, 'provider_id' | 'credential_id' | 'base_url'>> & {
|
||||
credential_id?: string | null
|
||||
base_url?: string | null
|
||||
}
|
||||
|
||||
export type RoutingCapability = 'embedding' | 'transcription' | 'speaker_matching'
|
||||
|
||||
export interface ModelBinding {
|
||||
provider_id: string
|
||||
model: string
|
||||
endpoint: string
|
||||
dimensions?: number | null
|
||||
}
|
||||
|
||||
export interface ModelRoutingConfig {
|
||||
version: number
|
||||
embedding: ModelBinding | null
|
||||
transcription: ModelBinding | null
|
||||
speaker_matching: ModelBinding | null
|
||||
}
|
||||
|
||||
export interface ModelRoutingResponse {
|
||||
config: ModelRoutingConfig
|
||||
local_backends: Array<{
|
||||
capability: RoutingCapability
|
||||
status: 'placeholder' | 'not_installed' | 'ready'
|
||||
message: string
|
||||
}>
|
||||
}
|
||||
|
||||
// ============ Tasks ============
|
||||
@@ -535,6 +568,51 @@ export interface OperationResponse {
|
||||
message?: string | null
|
||||
}
|
||||
|
||||
export type McpServerTransport = 'stdio' | 'streamable_http' | 'sse'
|
||||
export type McpServerState = 'stopped' | 'starting' | 'ready' | 'unhealthy' | 'error'
|
||||
|
||||
export interface McpServerInput {
|
||||
version?: number
|
||||
name: string
|
||||
transport: McpServerTransport
|
||||
command?: string | null
|
||||
args: string[]
|
||||
url?: string | null
|
||||
headers: Record<string, string>
|
||||
environment: Record<string, string>
|
||||
secret_environment_keys: string[]
|
||||
secret_header_keys: string[]
|
||||
permissions: string[]
|
||||
startup_timeout_seconds: number
|
||||
tool_timeout_seconds: number
|
||||
}
|
||||
|
||||
export interface McpServer extends Omit<McpServerInput, 'secret_environment_keys' | 'secret_header_keys'> {
|
||||
server_id: string
|
||||
version: number
|
||||
secret_environment: Record<string, boolean>
|
||||
secret_headers: Record<string, boolean>
|
||||
enabled: boolean
|
||||
trusted: boolean
|
||||
command_digest: string
|
||||
command_summary: string
|
||||
status: McpServerState
|
||||
tools_count: number
|
||||
protocol_version?: string | null
|
||||
remote_server_name?: string | null
|
||||
remote_server_version?: string | null
|
||||
error?: string | null
|
||||
last_tested_at?: string | null
|
||||
last_test_succeeded?: boolean | null
|
||||
}
|
||||
|
||||
export interface McpToolSummary {
|
||||
name: string
|
||||
remote_name: string
|
||||
description: string
|
||||
permission?: string | null
|
||||
}
|
||||
|
||||
export interface ApiNoteBlock {
|
||||
block_id: string
|
||||
note_id: string
|
||||
@@ -676,6 +754,9 @@ export interface ApiProviderPreset {
|
||||
base_url: string
|
||||
default_credential_id?: string | null
|
||||
requires_credential: boolean
|
||||
logo_id?: string
|
||||
description?: string
|
||||
capabilities?: string[]
|
||||
}
|
||||
|
||||
export interface ApiModelInfo {
|
||||
|
||||
@@ -27,6 +27,9 @@ beforeEach(() => {
|
||||
if (filePath === '/数据结构/红黑树.md') return '# 红黑树\n\n新的文件内容'
|
||||
throw new Error(`Unexpected file path: ${filePath}`)
|
||||
})
|
||||
vi.spyOn(workspaceService, 'getNoteId').mockImplementation(async (filePath) =>
|
||||
filePath.includes('红黑树') ? 'note-rbt' : 'note-welcome'
|
||||
)
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
|
||||
@@ -0,0 +1,169 @@
|
||||
// @vitest-environment happy-dom
|
||||
import { flushPromises, mount } from '@vue/test-utils'
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import type { McpServer } from '@/contracts'
|
||||
import * as service from '@/services/mcpServerService'
|
||||
import McpServersView from './McpServersView.vue'
|
||||
|
||||
vi.mock('@/services/mcpServerService', () => ({
|
||||
listMcpServers: vi.fn(), createMcpServer: vi.fn(), updateMcpServer: vi.fn(),
|
||||
deleteMcpServer: vi.fn(), trustMcpServer: vi.fn(), testMcpServer: vi.fn(),
|
||||
enableMcpServer: vi.fn(), disableMcpServer: vi.fn(), putMcpServerSecret: vi.fn(),
|
||||
}))
|
||||
|
||||
const server: McpServer = {
|
||||
server_id: 'server-1', version: 2, name: 'Remote', transport: 'streamable_http',
|
||||
command: null, args: [], url: 'https://mcp.example.test/mcp', headers: {}, environment: {},
|
||||
secret_environment: {}, secret_headers: { Authorization: false }, permissions: [],
|
||||
startup_timeout_seconds: 15, tool_timeout_seconds: 30, enabled: false, trusted: true,
|
||||
command_digest: 'a'.repeat(64), command_summary: 'https://mcp.example.test/mcp',
|
||||
status: 'stopped', tools_count: 1, last_test_succeeded: false,
|
||||
}
|
||||
|
||||
async function render(items: McpServer[] = []) {
|
||||
vi.mocked(service.listMcpServers).mockResolvedValue(items)
|
||||
const wrapper = mount(McpServersView, { global: { stubs: { AppIcon: true } } })
|
||||
await flushPromises()
|
||||
return wrapper
|
||||
}
|
||||
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks()
|
||||
vi.stubGlobal('confirm', vi.fn(() => true))
|
||||
})
|
||||
|
||||
describe('McpServersView', () => {
|
||||
it('switches transport templates and round-trips the JSON configuration mode', async () => {
|
||||
const wrapper = await render()
|
||||
await wrapper.findAll('button').find(button => button.text() === '新增服务器')!.trigger('click')
|
||||
await wrapper.findAll('button').find(button => button.text() === 'Streamable HTTP')!.trigger('click')
|
||||
expect(wrapper.find('input[placeholder="https://example.com/mcp"]').exists()).toBe(true)
|
||||
|
||||
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
|
||||
const raw = (wrapper.get('.json-editor').element as HTMLTextAreaElement).value
|
||||
expect(JSON.parse(raw)).toMatchObject({ transport: 'streamable_http', command: null })
|
||||
expect(raw).not.toContain('secret_value')
|
||||
|
||||
await wrapper.findAll('button').find(button => button.text() === '表单配置')!.trigger('click')
|
||||
expect(wrapper.text()).toContain('MCP URL')
|
||||
})
|
||||
|
||||
it('rejects invalid JSON without sending a create request', async () => {
|
||||
const wrapper = await render()
|
||||
await wrapper.findAll('button').find(button => button.text() === '新增服务器')!.trigger('click')
|
||||
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
|
||||
await wrapper.get('.json-editor').setValue('{invalid')
|
||||
await flushPromises()
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(wrapper.text()).toContain('服务器配置不是有效 JSON')
|
||||
expect(service.createMcpServer).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('keeps secrets request-only, exposes test failures, and confirms deletion', async () => {
|
||||
const wrapper = await render([server])
|
||||
const password = wrapper.get('input[type="password"]')
|
||||
await password.setValue('request-only-secret')
|
||||
vi.mocked(service.putMcpServerSecret).mockResolvedValue({} as never)
|
||||
await wrapper.findAll('button').find(button => button.text() === '保存')!.trigger('click')
|
||||
await flushPromises()
|
||||
expect(service.putMcpServerSecret).toHaveBeenCalledWith('server-1', 'Authorization', 'request-only-secret', 'header')
|
||||
expect((password.element as HTMLInputElement).value).toBe('')
|
||||
|
||||
vi.mocked(service.testMcpServer).mockRejectedValue(new Error('连接失败'))
|
||||
await wrapper.findAll('button').find(button => button.text().includes('测试连接'))!.trigger('click')
|
||||
await flushPromises()
|
||||
expect(wrapper.text()).toContain('连接失败')
|
||||
|
||||
vi.mocked(service.deleteMcpServer).mockResolvedValue({ status: 'completed' })
|
||||
await wrapper.findAll('button').find(button => button.text().includes('删除'))!.trigger('click')
|
||||
await flushPromises()
|
||||
expect(confirm).toHaveBeenCalled()
|
||||
expect(service.deleteMcpServer).toHaveBeenCalledWith('server-1')
|
||||
})
|
||||
|
||||
it('confirms permission changes before updating an existing server', async () => {
|
||||
const wrapper = await render([server])
|
||||
vi.mocked(service.updateMcpServer).mockResolvedValue(server)
|
||||
await wrapper.findAll('button').find(button => button.text().includes('编辑'))!.trigger('click')
|
||||
await wrapper.get('input[placeholder="network.request, notes.read"]').setValue('notes.read')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(confirm).toHaveBeenCalledWith(expect.stringContaining('旧测试与授权会失效'))
|
||||
expect(service.updateMcpServer).toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('saves an environment API key via the encrypted endpoint, not the config body', async () => {
|
||||
const wrapper = await render()
|
||||
vi.mocked(service.createMcpServer).mockResolvedValue({ ...server, server_id: 'new-server' })
|
||||
vi.mocked(service.putMcpServerSecret).mockResolvedValue({})
|
||||
await wrapper.findAll('button').find(button => button.text() === '新增服务器')!.trigger('click')
|
||||
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
|
||||
await wrapper.get('.json-editor').setValue(JSON.stringify({ command: 'uvx', environment: { MINIMAX_API_KEY: 'synthetic-only' }, secret_environment_keys: ['MINIMAX_API_KEY'] }))
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(service.createMcpServer).toHaveBeenCalledWith(expect.objectContaining({ environment: {}, secret_environment_keys: ['MINIMAX_API_KEY'] }))
|
||||
expect(JSON.stringify(vi.mocked(service.createMcpServer).mock.calls)).not.toContain('synthetic-only')
|
||||
expect(service.putMcpServerSecret).toHaveBeenCalledWith('new-server', 'MINIMAX_API_KEY', 'synthetic-only', 'environment')
|
||||
expect(wrapper.find('.modal-backdrop').exists()).toBe(false)
|
||||
})
|
||||
|
||||
it('retains imported keys over mode switches and retries partial saves without duplicates', async () => {
|
||||
const wrapper = await render()
|
||||
vi.mocked(service.createMcpServer).mockResolvedValue({ ...server, server_id: 'new-server', version: 1 })
|
||||
vi.mocked(service.updateMcpServer).mockResolvedValue({ ...server, server_id: 'new-server', version: 2 })
|
||||
vi.mocked(service.putMcpServerSecret).mockRejectedValueOnce(new Error('credential store unavailable')).mockResolvedValue({})
|
||||
await wrapper.findAll('button').find(button => button.text() === '新增服务器')!.trigger('click')
|
||||
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
|
||||
await wrapper.get('.json-editor').setValue(JSON.stringify({ command: 'uvx', env: { API_KEY: 'retry-value' } }))
|
||||
await wrapper.findAll('button').find(button => button.text() === '表单配置')!.trigger('click')
|
||||
expect(wrapper.text()).toContain('已识别 1 项密钥')
|
||||
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
|
||||
expect((wrapper.get('.json-editor').element as HTMLTextAreaElement).value).not.toContain('retry-value')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(wrapper.get('.modal-card [role="alert"]').text()).toContain('服务器配置已保存,但密钥保存失败')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(service.createMcpServer).toHaveBeenCalledTimes(1)
|
||||
expect(service.updateMcpServer).toHaveBeenCalledWith('new-server', expect.objectContaining({ version: 1 }))
|
||||
expect(service.putMcpServerSecret).toHaveBeenCalledTimes(2)
|
||||
expect(wrapper.find('.modal-backdrop').exists()).toBe(false)
|
||||
})
|
||||
|
||||
it('clears staged keys on cancel and accepts minimal JSON while editing', async () => {
|
||||
const wrapper = await render([server])
|
||||
await wrapper.findAll('button').find(button => button.text() === '新增服务器')!.trigger('click')
|
||||
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
|
||||
await wrapper.get('.json-editor').setValue('{"command":"uvx","env":{"API_KEY":"cancelled-value"}}')
|
||||
await wrapper.findAll('button').find(button => button.text() === '表单配置')!.trigger('click')
|
||||
await wrapper.findAll('button').find(button => button.text() === '取消')!.trigger('click')
|
||||
await wrapper.findAll('button').find(button => button.text().includes('编辑'))!.trigger('click')
|
||||
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
|
||||
await wrapper.get('.json-editor').setValue('{"name":"Minimal","url":"https://example.test/mcp"}')
|
||||
vi.mocked(service.updateMcpServer).mockResolvedValue(server)
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(service.updateMcpServer).toHaveBeenCalledWith('server-1', expect.objectContaining({ version: 2, headers: {}, args: [] }))
|
||||
expect(service.putMcpServerSecret).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('saves an imported Header secret after a case-only declaration rename', async () => {
|
||||
const wrapper = await render()
|
||||
vi.mocked(service.createMcpServer).mockResolvedValue({ ...server, secret_headers: { authorization: false } })
|
||||
vi.mocked(service.putMcpServerSecret).mockResolvedValue({})
|
||||
await wrapper.findAll('button').find(button => button.text() === '新增服务器')!.trigger('click')
|
||||
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
|
||||
await wrapper.get('.json-editor').setValue(JSON.stringify({ url: 'https://example.test/mcp', headers: { Authorization: 'synthetic-draft' } }))
|
||||
await wrapper.findAll('button').find(button => button.text() === '表单配置')!.trigger('click')
|
||||
await wrapper.get('textarea[placeholder="Authorization"]').setValue('authorization')
|
||||
await wrapper.findAll('button').find(button => button.text() === 'JSON 配置')!.trigger('click')
|
||||
expect(wrapper.text()).toContain('已识别 1 项密钥')
|
||||
expect((wrapper.get('.json-editor').element as HTMLTextAreaElement).value).not.toContain('synthetic-draft')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(service.createMcpServer).toHaveBeenCalledWith(expect.objectContaining({ headers: {}, secret_header_keys: ['authorization'] }))
|
||||
expect(service.putMcpServerSecret).toHaveBeenCalledWith('server-1', 'authorization', 'synthetic-draft', 'header')
|
||||
expect(wrapper.find('.modal-backdrop').exists()).toBe(false)
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,275 @@
|
||||
<script setup lang="ts">
|
||||
import { computed, onMounted, reactive, ref } from 'vue'
|
||||
import { Connection, Delete, EditPen, Plus, Refresh, VideoPlay } from '@element-plus/icons-vue'
|
||||
import AppIcon from '@/components/common/AppIcon.vue'
|
||||
import type { McpServer, McpServerInput, McpServerTransport } from '@/contracts'
|
||||
import * as service from '@/services/mcpServerService'
|
||||
import { emptyMcpConfig, mergeImportedSecrets, normalizeMcpConfig, parseMcpJson, type ImportedSecret, type SecretKind } from './configuration'
|
||||
|
||||
const servers = ref<McpServer[]>([])
|
||||
const busy = ref('')
|
||||
const error = ref('')
|
||||
const dialogOpen = ref(false)
|
||||
const editingId = ref<string | null>(null)
|
||||
const editingOriginal = ref<McpServer | null>(null)
|
||||
const editorMode = ref<'form' | 'json'>('form')
|
||||
const argsText = ref('')
|
||||
const environmentText = ref('{}')
|
||||
const headersText = ref('{}')
|
||||
const secretKeysText = ref('')
|
||||
const secretHeaderKeysText = ref('')
|
||||
const permissionsText = ref('')
|
||||
const rawConfig = ref('')
|
||||
const secretDrafts = reactive<Record<string, string>>({})
|
||||
const form = reactive<McpServerInput>(emptyMcpConfig())
|
||||
const importedSecrets = ref<ImportedSecret[]>([])
|
||||
|
||||
const dialogTitle = computed(() => editingId.value ? '编辑 MCP 服务器' : '新增 MCP 服务器')
|
||||
|
||||
async function load() {
|
||||
error.value = ''
|
||||
try { servers.value = await service.listMcpServers() }
|
||||
catch (cause) { error.value = message(cause, '读取 MCP 服务器失败') }
|
||||
}
|
||||
|
||||
function resetEditor(input: McpServerInput) {
|
||||
Object.assign(form, emptyMcpConfig(), { version: undefined }, input)
|
||||
argsText.value = input.args.join('\n')
|
||||
environmentText.value = JSON.stringify(input.environment, null, 2)
|
||||
headersText.value = JSON.stringify(input.headers, null, 2)
|
||||
secretKeysText.value = input.secret_environment_keys.join('\n')
|
||||
secretHeaderKeysText.value = input.secret_header_keys.join('\n')
|
||||
permissionsText.value = input.permissions.join(', ')
|
||||
editorMode.value = 'form'
|
||||
rawConfig.value = ''
|
||||
}
|
||||
|
||||
function openCreate() {
|
||||
if (busy.value) return
|
||||
error.value = ''
|
||||
importedSecrets.value = []
|
||||
editingId.value = null
|
||||
editingOriginal.value = null
|
||||
resetEditor(emptyMcpConfig())
|
||||
dialogOpen.value = true
|
||||
}
|
||||
|
||||
function openEdit(server: McpServer) {
|
||||
if (busy.value) return
|
||||
error.value = ''
|
||||
importedSecrets.value = []
|
||||
editingId.value = server.server_id
|
||||
editingOriginal.value = server
|
||||
resetEditor({
|
||||
version: server.version, name: server.name, transport: server.transport,
|
||||
command: server.command, args: [...server.args], url: server.url,
|
||||
headers: { ...server.headers }, environment: { ...server.environment },
|
||||
secret_environment_keys: Object.keys(server.secret_environment),
|
||||
secret_header_keys: Object.keys(server.secret_headers), permissions: [...server.permissions],
|
||||
startup_timeout_seconds: server.startup_timeout_seconds,
|
||||
tool_timeout_seconds: server.tool_timeout_seconds,
|
||||
})
|
||||
dialogOpen.value = true
|
||||
}
|
||||
|
||||
function applyTemplate(transport: McpServerTransport) {
|
||||
form.transport = transport
|
||||
if (transport === 'stdio') {
|
||||
form.command = 'uvx'; form.url = null
|
||||
argsText.value = '--isolated\n--from\npackage-name==1.0.0\nserver-command'
|
||||
} else {
|
||||
form.command = null; argsText.value = ''; form.url = transport === 'sse' ? 'http://127.0.0.1:3000/sse' : 'http://127.0.0.1:3000/mcp'
|
||||
}
|
||||
}
|
||||
|
||||
function parseObject(value: string, label: string): Record<string, string> {
|
||||
let parsed: unknown
|
||||
try { parsed = JSON.parse(value || '{}') } catch { throw new Error(`${label}必须是 JSON 对象`) }
|
||||
if (!parsed || Array.isArray(parsed) || typeof parsed !== 'object' || Object.values(parsed).some(item => typeof item !== 'string')) throw new Error(`${label}必须是字符串键值 JSON 对象`)
|
||||
return parsed as Record<string, string>
|
||||
}
|
||||
|
||||
function formPayload(): McpServerInput {
|
||||
const stdio = form.transport === 'stdio'
|
||||
return {
|
||||
version: form.version,
|
||||
name: form.name.trim(), transport: form.transport,
|
||||
command: stdio ? form.command?.trim() : null,
|
||||
args: stdio ? argsText.value.split('\n').map(value => value.trim()).filter(Boolean) : [],
|
||||
url: stdio ? null : form.url?.trim(),
|
||||
headers: stdio ? {} : parseObject(headersText.value, '普通 Header'),
|
||||
environment: stdio ? parseObject(environmentText.value, '普通环境变量') : {},
|
||||
secret_environment_keys: stdio ? splitKeys(secretKeysText.value) : [],
|
||||
secret_header_keys: stdio ? [] : splitKeys(secretHeaderKeysText.value),
|
||||
permissions: permissionsText.value.split(',').map(value => value.trim()).filter(Boolean),
|
||||
startup_timeout_seconds: form.startup_timeout_seconds,
|
||||
tool_timeout_seconds: form.tool_timeout_seconds,
|
||||
}
|
||||
}
|
||||
|
||||
function payload(requireConnection = true): McpServerInput {
|
||||
const { config, secrets } = editorMode.value === 'form'
|
||||
? normalizeMcpConfig(formPayload(), '', requireConnection) : parseMcpJson(rawConfig.value, form.name, requireConnection)
|
||||
// Keep only still-declared drafts. A mode switch must not discard imported keys,
|
||||
// and editing the declaration must not later send a removed key to the Secret API.
|
||||
importedSecrets.value = mergeImportedSecrets(config, importedSecrets.value, secrets)
|
||||
if (editingId.value) config.version = form.version
|
||||
if (editorMode.value === 'json') rawConfig.value = JSON.stringify(config, null, 2)
|
||||
else {
|
||||
environmentText.value = JSON.stringify(config.environment, null, 2)
|
||||
headersText.value = JSON.stringify(config.headers, null, 2)
|
||||
secretKeysText.value = config.secret_environment_keys.join('\n')
|
||||
secretHeaderKeysText.value = config.secret_header_keys.join('\n')
|
||||
}
|
||||
return config
|
||||
}
|
||||
|
||||
function switchMode(mode: 'form' | 'json') {
|
||||
try {
|
||||
if (mode === editorMode.value) return
|
||||
error.value = ''
|
||||
if (mode === 'json') rawConfig.value = JSON.stringify(payload(false), null, 2)
|
||||
else resetEditor(payload(false))
|
||||
editorMode.value = mode
|
||||
} catch (cause) { error.value = message(cause, '配置转换失败') }
|
||||
}
|
||||
|
||||
async function save() {
|
||||
if (busy.value) return
|
||||
let saved: McpServer | undefined
|
||||
try {
|
||||
error.value = ''
|
||||
const input = payload()
|
||||
if (!input.name || (input.transport === 'stdio' ? !input.command : !input.url)) throw new Error('请填写服务器名称和连接地址')
|
||||
if (editingOriginal.value && executionChanged(editingOriginal.value, input) && !confirm('连接命令、地址或认证配置已变化,保存后旧测试与授权会失效。是否保存?')) return
|
||||
busy.value = 'save'
|
||||
saved = editingId.value ? await service.updateMcpServer(editingId.value, input) : await service.createMcpServer(input)
|
||||
// Commit the returned ID/version before saving secrets so a partial failure can
|
||||
// retry this server instead of creating a duplicate or sending a stale version.
|
||||
editingId.value = saved.server_id
|
||||
editingOriginal.value = saved
|
||||
resetEditor({ ...input, version: saved.version })
|
||||
for (const item of [...importedSecrets.value]) {
|
||||
await service.putMcpServerSecret(saved.server_id, item.key, item.value, item.kind)
|
||||
importedSecrets.value = importedSecrets.value.filter(candidate => candidate !== item)
|
||||
}
|
||||
closeEditor()
|
||||
await load()
|
||||
} catch (cause) {
|
||||
if (saved) await load()
|
||||
error.value = `${saved ? '服务器配置已保存,但密钥保存失败;可点击保存重试。' : ''}${message(cause, '保存失败')}`
|
||||
}
|
||||
finally { busy.value = '' }
|
||||
}
|
||||
|
||||
function closeEditor() {
|
||||
importedSecrets.value = []
|
||||
rawConfig.value = ''
|
||||
environmentText.value = '{}'
|
||||
headersText.value = '{}'
|
||||
dialogOpen.value = false
|
||||
}
|
||||
|
||||
function executionChanged(server: McpServer, input: McpServerInput) {
|
||||
const sortedEntries = (value: Record<string, string>) => Object.entries(value).sort(([left], [right]) => left.localeCompare(right))
|
||||
const current = [
|
||||
server.transport, server.command, server.args, server.url,
|
||||
sortedEntries(server.headers), sortedEntries(server.environment),
|
||||
Object.keys(server.secret_headers).sort(), Object.keys(server.secret_environment).sort(),
|
||||
[...server.permissions].sort(), server.startup_timeout_seconds, server.tool_timeout_seconds,
|
||||
]
|
||||
const next = [
|
||||
input.transport, input.command, input.args, input.url,
|
||||
sortedEntries(input.headers), sortedEntries(input.environment),
|
||||
[...input.secret_header_keys].sort(), [...input.secret_environment_keys].sort(),
|
||||
[...input.permissions].sort(), input.startup_timeout_seconds, input.tool_timeout_seconds,
|
||||
]
|
||||
return JSON.stringify(current) !== JSON.stringify(next)
|
||||
}
|
||||
|
||||
async function approve(server: McpServer): Promise<McpServer | null> {
|
||||
if (server.trusted) return server
|
||||
const localWarning = server.transport === 'stdio' ? '\n\n本机进程尚无系统级沙箱,仅应运行可信服务器。' : '\n\n连接可能向该地址发送配置的 Header。'
|
||||
if (!confirm(`请确认 MCP 连接:\n\n${server.command_summary}${localWarning}\n\n是否继续?`)) return null
|
||||
return service.trustMcpServer(server)
|
||||
}
|
||||
|
||||
async function test(server: McpServer) { await act(server, 'test', current => service.testMcpServer(current.server_id)) }
|
||||
async function toggle(server: McpServer) { await act(server, 'toggle', current => current.enabled ? service.disableMcpServer(current.server_id) : service.enableMcpServer(current.server_id)) }
|
||||
async function act(server: McpServer, action: string, operation: (server: McpServer) => Promise<McpServer>) {
|
||||
busy.value = `${action}:${server.server_id}`; error.value = ''
|
||||
try { const current = action === 'toggle' && server.enabled ? server : await approve(server); if (!current) return; await operation(current); await load() }
|
||||
catch (cause) { error.value = message(cause, '操作失败') }
|
||||
finally { busy.value = '' }
|
||||
}
|
||||
|
||||
async function remove(server: McpServer) {
|
||||
if (!confirm(`删除“${server.name}”及其加密凭据?`)) return
|
||||
try { busy.value = `delete:${server.server_id}`; await service.deleteMcpServer(server.server_id); await load() }
|
||||
catch (cause) { error.value = message(cause, '删除失败') } finally { busy.value = '' }
|
||||
}
|
||||
|
||||
async function saveSecret(server: McpServer, key: string, kind: SecretKind) {
|
||||
const draftKey = `${server.server_id}:${kind}:${key}`
|
||||
const value = secretDrafts[draftKey]?.trim()
|
||||
if (!value) return
|
||||
try { busy.value = `secret:${draftKey}`; await service.putMcpServerSecret(server.server_id, key, value, kind); secretDrafts[draftKey] = ''; await load() }
|
||||
catch (cause) { error.value = message(cause, '保存密钥失败') } finally { busy.value = '' }
|
||||
}
|
||||
|
||||
function splitKeys(value: string) { return value.split(/[\n,]/).map(item => item.trim()).filter(Boolean) }
|
||||
function message(cause: unknown, fallback: string) { return cause instanceof Error ? cause.message : fallback }
|
||||
onMounted(load)
|
||||
</script>
|
||||
|
||||
<template>
|
||||
<section class="feature-page mcp-page">
|
||||
<header class="feature-header"><div><h1>MCP 服务器</h1><p>管理独立 MCP Server 的连接、凭据与工具生命周期。</p></div><div class="inline-actions"><button class="button-secondary" :disabled="!!busy" @click="load"><AppIcon :icon="Refresh" /> 刷新</button><button class="button-primary" @click="openCreate"><AppIcon :icon="Plus" /> 新增服务器</button></div></header>
|
||||
<div class="notice-banner">stdio 本机进程仅在开发环境开放;Streamable HTTP 为首选远程传输,SSE 仅用于兼容旧服务器。uvx 隔离依赖但不是安全沙箱。</div>
|
||||
<div v-if="error" class="error-banner">{{ error }}</div>
|
||||
<div v-if="!servers.length" class="panel empty"><AppIcon :icon="Connection" :size="34" /><h2>尚未配置 MCP 服务器</h2><p>添加 Server,测试连接成功后才能启用工具。</p><button class="button-primary" @click="openCreate">新增服务器</button></div>
|
||||
<div v-else class="server-list">
|
||||
<article v-for="server in servers" :key="server.server_id" class="panel server-card">
|
||||
<div class="server-main"><div class="server-title"><AppIcon :icon="Connection" :size="24" /><div><h2>{{ server.name }}</h2><code>{{ server.command_summary }}</code></div></div><span class="badge" :class="{ success: server.status === 'ready', error: ['error','unhealthy'].includes(server.status) }">{{ server.status }}</span></div>
|
||||
<div class="metadata"><span>{{ server.transport }}</span><span>v{{ server.version }}</span><span>{{ server.tools_count }} 个工具</span><span>{{ server.trusted ? '连接已确认' : '等待确认连接' }}</span><span v-if="server.last_test_succeeded">当前配置测试成功</span><span v-if="server.remote_server_name">{{ server.remote_server_name }} {{ server.remote_server_version }}</span></div>
|
||||
<div v-if="server.error" class="error-banner compact">{{ server.error }}</div>
|
||||
<div v-if="Object.keys(server.secret_environment).length || Object.keys(server.secret_headers).length" class="secrets">
|
||||
<label v-for="(configured, key) in server.secret_environment" :key="`env:${key}`"><span>环境变量 · {{ key }} <small>{{ configured ? '已加密保存' : '未配置' }}</small></span><span class="secret-input"><input v-model="secretDrafts[`${server.server_id}:environment:${key}`]" type="password" autocomplete="new-password" placeholder="输入后保存(不会回显)"><button class="button-secondary" @click="saveSecret(server, key, 'environment')">保存</button></span></label>
|
||||
<label v-for="(configured, key) in server.secret_headers" :key="`header:${key}`"><span>HTTP Header · {{ key }} <small>{{ configured ? '已加密保存' : '未配置' }}</small></span><span class="secret-input"><input v-model="secretDrafts[`${server.server_id}:header:${key}`]" type="password" autocomplete="new-password" placeholder="输入后保存(不会回显)"><button class="button-secondary" @click="saveSecret(server, key, 'header')">保存</button></span></label>
|
||||
</div>
|
||||
<footer class="card-actions"><button class="button-secondary" :disabled="!!busy || server.enabled" @click="test(server)"><AppIcon :icon="VideoPlay" /> 测试连接</button><button class="button-secondary" :disabled="!!busy" @click="openEdit(server)"><AppIcon :icon="EditPen" /> 编辑</button><button class="button-danger" :disabled="!!busy" @click="remove(server)"><AppIcon :icon="Delete" /> 删除</button><button class="button-primary" :disabled="!!busy || (!server.enabled && !server.last_test_succeeded)" :title="!server.enabled && !server.last_test_succeeded ? '请先测试当前配置' : ''" @click="toggle(server)">{{ server.enabled ? '停用' : '启用' }}</button></footer>
|
||||
</article>
|
||||
</div>
|
||||
|
||||
<div v-if="dialogOpen" class="modal-backdrop" @click.self="!busy && closeEditor()">
|
||||
<form class="modal-card" @submit.prevent="save">
|
||||
<fieldset :disabled="!!busy" class="editor-fields">
|
||||
<header><h2><AppIcon :icon="Plus" /> {{ dialogTitle }}</h2><button type="button" class="close" @click="closeEditor">×</button></header>
|
||||
<div v-if="error" class="error-banner" role="alert">{{ error }}</div>
|
||||
<div v-if="importedSecrets.length" class="notice-banner">已识别 {{ importedSecrets.length }} 项密钥,保存时将单独加密,不会写入普通服务器配置;取消将清除未保存密钥。</div>
|
||||
<div class="mode-tabs"><button type="button" :class="{ active: editorMode === 'form' }" @click="switchMode('form')">表单配置</button><button type="button" :class="{ active: editorMode === 'json' }" @click="switchMode('json')">JSON 配置</button></div>
|
||||
<template v-if="editorMode === 'form'">
|
||||
<label>服务器名称<input v-model="form.name" maxlength="80" placeholder="例如:文件系统工具"></label>
|
||||
<div class="template-row"><span>服务器配置</span><button type="button" class="template" :class="{ active: form.transport === 'stdio' }" @click="applyTemplate('stdio')">stdio 模板</button><button type="button" class="template" :class="{ active: form.transport === 'streamable_http' }" @click="applyTemplate('streamable_http')">Streamable HTTP</button><button type="button" class="template" :class="{ active: form.transport === 'sse' }" @click="applyTemplate('sse')">SSE(兼容)</button></div>
|
||||
<template v-if="form.transport === 'stdio'"><label>可执行命令<input v-model="form.command" placeholder="uvx、npx 或可信可执行文件路径"></label><label>参数(每行一项)<textarea v-model="argsText" rows="5"></textarea></label><div class="two-columns"><label>普通环境变量(JSON)<textarea v-model="environmentText" rows="5"></textarea></label><label>敏感环境变量名(每行一项)<textarea v-model="secretKeysText" rows="5" placeholder="API_KEY"></textarea></label></div></template>
|
||||
<template v-else><label>MCP URL<input v-model="form.url" placeholder="https://example.com/mcp"></label><div class="two-columns"><label>普通 Header(JSON)<textarea v-model="headersText" rows="5" placeholder='{"X-Client":"NotesAgent"}'></textarea></label><label>敏感 Header 名(每行一项)<textarea v-model="secretHeaderKeysText" rows="5" placeholder="Authorization"></textarea></label></div></template>
|
||||
<label>声明权限(逗号分隔,可选)<input v-model="permissionsText" placeholder="network.request, notes.read"></label>
|
||||
<div class="two-columns"><label>启动超时(秒)<input v-model.number="form.startup_timeout_seconds" type="number" min="1" max="120"></label><label>工具超时(秒)<input v-model.number="form.tool_timeout_seconds" type="number" min="1" max="300"></label></div>
|
||||
</template>
|
||||
<label v-else>服务器 JSON 配置<textarea v-model="rawConfig" class="json-editor" rows="22" spellcheck="false"></textarea><small>支持 NotesAgent 配置、command/args/env 和单服务器 mcpServers 配置。已声明的 Secret 及常见 API Key、Token、Authorization 会拆分后加密保存。其他敏感值请显式声明;不要把密钥放入命令或参数。</small><small>兼容导入 timeout 为启动超时,sse_read_timeout 为工具等待上限(不保留 SSE 读取超时语义)。</small></label>
|
||||
<footer><button type="button" class="button-secondary" @click="closeEditor">取消</button><button class="button-primary" :disabled="busy === 'save'">保存</button></footer>
|
||||
</fieldset>
|
||||
</form>
|
||||
</div>
|
||||
</section>
|
||||
</template>
|
||||
|
||||
<style scoped>
|
||||
.editor-fields { display: grid; gap: var(--space-lg); border: 0; padding: 0; margin: 0; min-width: 0; }
|
||||
.mcp-page { overflow: auto; }.notice-banner,.error-banner { margin-bottom: var(--space-lg); }.server-list { display: grid; gap: var(--space-lg); }.server-card { display: grid; gap: var(--space-md); }
|
||||
.server-main,.server-title,.metadata,.card-actions,.inline-actions,.template-row,.modal-card header,.modal-card footer { display: flex; align-items: center; gap: var(--space-sm); }.server-main { justify-content: space-between; }.server-title { align-items: flex-start; }.server-title h2 { margin-bottom: 4px; }.server-title code { color: var(--color-text-secondary); overflow-wrap: anywhere; }.metadata { flex-wrap: wrap; color: var(--color-text-tertiary); font-size: var(--font-size-sm); }.metadata span + span::before { content: '·'; margin-right: var(--space-sm); }.compact { margin: 0; }
|
||||
.card-actions { justify-content: flex-end; border-top: 1px solid var(--color-border-subtle); padding-top: var(--space-md); }.empty { text-align: center; place-items: center; display: grid; gap: var(--space-md); padding: 64px; }.secrets { border: 1px solid var(--color-border-subtle); border-radius: var(--radius-md); padding: var(--space-md); display: grid; gap: var(--space-sm); }.secrets label { display: grid; grid-template-columns: minmax(220px,.7fr) 1fr; align-items: center; gap: var(--space-md); }.secrets small,.modal-card small { color: var(--color-text-tertiary); }.secret-input { display: flex; gap: var(--space-sm); }.secret-input input { flex: 1; }
|
||||
.modal-backdrop { position: fixed; inset: 0; z-index: 1000; background: rgb(0 0 0 / .48); display: grid; place-items: center; padding: var(--space-xl); }.modal-card { width: min(800px,100%); max-height: calc(100vh - 48px); overflow: auto; background: var(--color-background-primary); border: 1px solid var(--color-border-default); border-radius: var(--radius-xl); box-shadow: var(--shadow-xl); padding: var(--space-xl); display: grid; gap: var(--space-lg); animation: modal-in var(--motion-normal) ease-out; }.modal-card header,.modal-card footer { justify-content: space-between; }.modal-card footer { justify-content: flex-end; }.modal-card label { display: grid; gap: var(--space-xs); font-weight: 600; }.modal-card input,.modal-card textarea { width: 100%; border: 1px solid var(--color-border-default); border-radius: var(--radius-md); padding: 10px 12px; color: var(--color-text-primary); background: var(--color-background-secondary); font: inherit; }.modal-card textarea { resize: vertical; font-family: var(--font-family-mono); font-size: var(--font-size-sm); }.json-editor { line-height: 1.55; }.close { border: 0; background: transparent; color: var(--color-text-secondary); font-size: 28px; cursor: pointer; }
|
||||
.template-row { flex-wrap: wrap; }.template-row > span { margin-right: auto; font-weight: 600; }.template,.mode-tabs button { border: 1px solid var(--color-border-default); background: var(--color-background-secondary); color: var(--color-text-secondary); padding: 7px 10px; border-radius: var(--radius-md); cursor: pointer; }.template.active,.mode-tabs button.active { color: var(--color-accent-primary); border-color: var(--color-accent-primary); background: var(--color-accent-soft); }.mode-tabs { display: inline-flex; justify-self: start; gap: 2px; padding: 3px; border-radius: var(--radius-md); background: var(--color-background-secondary); }.two-columns { display: grid; grid-template-columns: 1fr 1fr; gap: var(--space-md); }
|
||||
@keyframes modal-in { from { opacity: 0; transform: translateY(8px) scale(.99); } } @media (max-width:720px) { .two-columns,.secrets label { grid-template-columns:1fr; }.card-actions { flex-wrap:wrap; } }
|
||||
</style>
|
||||
@@ -0,0 +1,63 @@
|
||||
import { describe, expect, it } from 'vitest'
|
||||
import { emptyMcpConfig, mergeImportedSecrets, normalizeMcpConfig, parseMcpJson } from './configuration'
|
||||
|
||||
describe('MCP configuration normalization', () => {
|
||||
it('retains renamed HTTP drafts with the latest spelling and value', () => {
|
||||
const config = { ...emptyMcpConfig(), secret_header_keys: ['authorization'] }
|
||||
const previous = [{ kind: 'header' as const, key: 'Authorization', value: 'old-value' }]
|
||||
expect(mergeImportedSecrets(config, previous, [])).toEqual([{ kind: 'header', key: 'authorization', value: 'old-value' }])
|
||||
expect(mergeImportedSecrets(config, previous, [{ kind: 'header', key: 'AUTHORIZATION', value: 'new-value' }])).toEqual([{ kind: 'header', key: 'authorization', value: 'new-value' }])
|
||||
expect(mergeImportedSecrets(emptyMcpConfig(), previous, [])).toEqual([])
|
||||
})
|
||||
|
||||
it('does not transfer an environment draft across a case-only rename', () => {
|
||||
const config = { ...emptyMcpConfig(), secret_environment_keys: ['TOKEN', 'token'] }
|
||||
const previous = [{ kind: 'environment' as const, key: 'TOKEN', value: 'upper' }, { kind: 'environment' as const, key: 'token', value: 'lower' }]
|
||||
expect(mergeImportedSecrets(config, previous, [])).toEqual(previous)
|
||||
expect(mergeImportedSecrets({ ...config, secret_environment_keys: ['token'] }, [previous[0]!], [])).toEqual([])
|
||||
})
|
||||
it('fills backend defaults for minimal JSON', () => {
|
||||
const { config } = parseMcpJson('{"name":"demo","command":"uvx"}')
|
||||
expect(config).toMatchObject({ transport: 'stdio', args: [], headers: {}, environment: {}, permissions: [], secret_header_keys: [] })
|
||||
})
|
||||
|
||||
it('extracts a key pasted into environment despite its existing secret declaration', () => {
|
||||
const { config, secrets } = normalizeMcpConfig({
|
||||
name: 'MiniMax', command: 'uvx', secret_environment_keys: ['MINIMAX_API_KEY'],
|
||||
environment: { MINIMAX_API_KEY: 'synthetic-key', MINIMAX_API_HOST: 'https://api.minimaxi.com' },
|
||||
})
|
||||
expect(config.environment).toEqual({ MINIMAX_API_HOST: 'https://api.minimaxi.com' })
|
||||
expect(config.secret_environment_keys).toEqual(['MINIMAX_API_KEY'])
|
||||
expect(JSON.stringify(config)).not.toContain('synthetic-key')
|
||||
expect(secrets).toEqual([{ kind: 'environment', key: 'MINIMAX_API_KEY', value: 'synthetic-key' }])
|
||||
})
|
||||
|
||||
it('imports a standard single-server wrapper and legacy timeouts', () => {
|
||||
const { config, secrets } = normalizeMcpConfig({ mcpServers: { MiniMax: {
|
||||
command: 'uvx', args: ['--with', 'mcp<2', 'minimax-coding-plan-mcp', '-y'],
|
||||
env: { MINIMAX_API_KEY: 'synthetic-key' }, timeout: 120, sse_read_timeout: 300,
|
||||
} } })
|
||||
expect(config).toMatchObject({ name: 'MiniMax', transport: 'stdio', environment: {}, startup_timeout_seconds: 120, tool_timeout_seconds: 300 })
|
||||
expect(secrets).toHaveLength(1)
|
||||
})
|
||||
|
||||
it('extracts case-insensitive HTTP credentials without duplicate declarations', () => {
|
||||
const { config, secrets } = normalizeMcpConfig({ url: 'https://example.test/mcp', headers: { authorization: 'synthetic' }, secret_header_keys: ['Authorization'] })
|
||||
expect(config.headers).toEqual({})
|
||||
expect(config.secret_header_keys).toEqual(['Authorization'])
|
||||
expect(secrets[0]?.key).toBe('Authorization')
|
||||
})
|
||||
|
||||
it.each([
|
||||
[{ command: 'uvx', args: 'not-array' }, 'args'],
|
||||
[{ command: 'uvx', environment: [] }, 'environment'],
|
||||
[{ command: 'uvx', timeout: 121 }, '启动超时'],
|
||||
[{ command: 'uvx', args: ['[https://example.test](https://example.test)'] }, '纯 URL'],
|
||||
[{ command: 'uvx', api_key: 'do-not-echo' }, '顶层'],
|
||||
[{ command: 'uvx', env: {}, environment: {} }, '只保留一个'],
|
||||
[{ mcpServers: { one: {}, two: {} } }, '一次导入一个'],
|
||||
])('rejects invalid fields without leaking their values', (input, hint) => {
|
||||
expect(() => normalizeMcpConfig(input)).toThrow(hint)
|
||||
try { normalizeMcpConfig(input) } catch (error) { expect(String(error)).not.toContain('do-not-echo') }
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,139 @@
|
||||
import type { McpServerInput } from '@/contracts'
|
||||
|
||||
export type SecretKind = 'environment' | 'header'
|
||||
export interface ImportedSecret { kind: SecretKind; key: string; value: string }
|
||||
|
||||
export function mergeImportedSecrets(config: McpServerInput, previous: ImportedSecret[], incoming: ImportedSecret[]): ImportedSecret[] {
|
||||
const merged = new Map<string, ImportedSecret>()
|
||||
for (const item of [...previous, ...incoming]) {
|
||||
const normalize = (key: string) => item.kind === 'header' ? key.toLowerCase() : key
|
||||
const keys = item.kind === 'header' ? config.secret_header_keys : config.secret_environment_keys
|
||||
const declared = keys.find(key => normalize(key) === normalize(item.key))
|
||||
if (declared === undefined) continue
|
||||
// HTTP identity is case-insensitive, but the Secret API requires the current
|
||||
// declared spelling. New inline values replace older drafts of that identity.
|
||||
merged.set(`${item.kind}:${normalize(declared)}`, { ...item, key: declared })
|
||||
}
|
||||
return [...merged.values()]
|
||||
}
|
||||
|
||||
export function emptyMcpConfig(): McpServerInput {
|
||||
return {
|
||||
name: '', transport: 'stdio', command: '', args: [], url: null, headers: {},
|
||||
environment: {}, secret_environment_keys: [], secret_header_keys: [], permissions: [],
|
||||
startup_timeout_seconds: 15, tool_timeout_seconds: 30,
|
||||
}
|
||||
}
|
||||
|
||||
function object(value: unknown, label: string): Record<string, unknown> {
|
||||
if (!value || Array.isArray(value) || typeof value !== 'object') throw new Error(`${label}必须是 JSON 对象`)
|
||||
return value as Record<string, unknown>
|
||||
}
|
||||
|
||||
function strings(value: unknown, label: string): string[] {
|
||||
if (value === undefined) return []
|
||||
if (!Array.isArray(value) || value.some(item => typeof item !== 'string')) throw new Error(`${label}必须是字符串数组`)
|
||||
return [...value]
|
||||
}
|
||||
|
||||
function entries(value: unknown, label: string): Record<string, string> {
|
||||
if (value === undefined) return {}
|
||||
const result = object(value, label)
|
||||
if (Object.values(result).some(item => typeof item !== 'string')) throw new Error(`${label}必须是字符串键值 JSON 对象`)
|
||||
return { ...result } as Record<string, string>
|
||||
}
|
||||
|
||||
function timeout(value: unknown, fallback: number, max: number, label: string): number {
|
||||
if (value === undefined) return fallback
|
||||
if (typeof value !== 'number' || !Number.isFinite(value) || value < 1 || value > max) throw new Error(`${label}必须是 1–${max} 秒之间的数字`)
|
||||
return value
|
||||
}
|
||||
|
||||
// Do not silently rewrite executable arguments or secret values copied from chat.
|
||||
function checkUrl(value: string, label: string) {
|
||||
if (/^\[https?:\/\//i.test(value)) throw new Error(`${label}请填写纯 URL,不要粘贴 Markdown 链接`)
|
||||
}
|
||||
|
||||
export function parseMcpJson(raw: string, fallbackName = '', requireConnection = true) {
|
||||
let parsed: unknown
|
||||
try { parsed = JSON.parse(raw) }
|
||||
catch { throw new Error('服务器配置不是有效 JSON;请检查逗号、引号和无效的 \\_ 转义') }
|
||||
return normalizeMcpConfig(parsed, fallbackName, requireConnection)
|
||||
}
|
||||
|
||||
/** Normalize external client JSON before it reaches either the form or the API.
|
||||
* Inline secrets leave the public config here and are sent only to the Secret API.
|
||||
*/
|
||||
export function normalizeMcpConfig(parsed: unknown, fallbackName = '', requireConnection = true) {
|
||||
let raw = object(parsed, '服务器配置')
|
||||
if ('mcpServers' in raw) {
|
||||
const servers = Object.entries(object(raw.mcpServers, 'mcpServers'))
|
||||
if (servers.length !== 1) throw new Error('请一次导入一个 MCP 服务器')
|
||||
fallbackName = servers[0]![0]
|
||||
raw = object(servers[0]![1], '服务器配置')
|
||||
}
|
||||
const allowed = new Set([...Object.keys(emptyMcpConfig()), 'version', 'env', 'type', 'timeout', 'sse_read_timeout'])
|
||||
if (Object.keys(raw).some(key => !allowed.has(key))) {
|
||||
// Never echo arbitrary unknown keys: pasted secrets sometimes become JSON keys.
|
||||
throw new Error('服务器配置含不支持的字段;API Key 请放在 env/environment 的对应变量中,不要放在顶层')
|
||||
}
|
||||
if (raw.env !== undefined && raw.environment !== undefined) throw new Error('env 与 environment 请只保留一个,避免覆盖配置')
|
||||
const transport = raw.transport ?? raw.type ?? (raw.url ? 'streamable_http' : 'stdio')
|
||||
if (!['stdio', 'streamable_http', 'sse'].includes(transport as string)) throw new Error('transport 必须是 stdio、streamable_http 或 sse')
|
||||
const config = emptyMcpConfig()
|
||||
config.transport = transport as McpServerInput['transport']
|
||||
const name = raw.name ?? (fallbackName || (typeof raw.command === 'string' ? raw.command : 'MCP 服务器'))
|
||||
if (typeof name !== 'string' || (requireConnection && !name.trim()) || name.trim().length > 80) throw new Error('服务器名称必须为 1–80 个字符')
|
||||
config.name = name.trim()
|
||||
for (const key of ['command', 'url'] as const) {
|
||||
const value = raw[key]
|
||||
if (value !== undefined && value !== null && typeof value !== 'string') throw new Error(`${key}必须是字符串`)
|
||||
config[key] = typeof value === 'string' ? value.trim() : null
|
||||
}
|
||||
config.args = strings(raw.args, 'args')
|
||||
if (config.args.length > 64) throw new Error('args 最多允许 64 项')
|
||||
for (const value of config.args) checkUrl(value, 'args 中的地址')
|
||||
config.environment = entries(raw.environment ?? raw.env, 'environment/env')
|
||||
config.headers = entries(raw.headers, 'headers')
|
||||
config.secret_environment_keys = [...new Set(strings(raw.secret_environment_keys, 'secret_environment_keys'))]
|
||||
config.secret_header_keys = [...new Set(strings(raw.secret_header_keys, 'secret_header_keys'))]
|
||||
config.permissions = strings(raw.permissions, 'permissions')
|
||||
config.startup_timeout_seconds = timeout(raw.startup_timeout_seconds ?? raw.timeout, 15, 120, '启动超时')
|
||||
// Compatibility policy: legacy read timeout becomes the tool wait budget, not an SSE transport setting.
|
||||
config.tool_timeout_seconds = timeout(raw.tool_timeout_seconds ?? raw.sse_read_timeout, 30, 300, '工具超时')
|
||||
if (config.transport === 'stdio') {
|
||||
if (requireConnection && !config.command) throw new Error('stdio 配置必须填写 command')
|
||||
if (config.url || Object.keys(config.headers).length || config.secret_header_keys.length) throw new Error('stdio 配置不能包含 URL 或 HTTP Header')
|
||||
} else {
|
||||
if (requireConnection && !config.url) throw new Error('HTTP/SSE 配置必须填写 url')
|
||||
if (config.url) {
|
||||
checkUrl(config.url, 'url')
|
||||
let url: URL
|
||||
try { url = new URL(config.url) } catch { throw new Error('url 必须是有效的 HTTP(S) 地址') }
|
||||
if (!['http:', 'https:'].includes(url.protocol) || url.username || url.password || url.hash) throw new Error('url 必须为不含账号密码或片段的 HTTP(S) 地址')
|
||||
}
|
||||
if (config.command || config.args.length || Object.keys(config.environment).length || config.secret_environment_keys.length) throw new Error('HTTP/SSE 配置不能包含 command、args 或环境变量')
|
||||
}
|
||||
const secrets: ImportedSecret[] = []
|
||||
for (const kind of ['environment', 'header'] as const) {
|
||||
const values = kind === 'environment' ? config.environment : config.headers
|
||||
const keys = kind === 'environment' ? config.secret_environment_keys : config.secret_header_keys
|
||||
const identity = (key: string) => kind === 'header' ? key.toLowerCase() : key
|
||||
const allKeys = [...Object.keys(values), ...keys]
|
||||
if (kind === 'header' && (new Set(keys.map(identity)).size !== keys.length || new Set(Object.keys(values).map(identity)).size !== Object.keys(values).length)) throw new Error('HTTP Header 名称不能仅大小写不同而重复声明')
|
||||
const validKey = kind === 'environment' ? /^[A-Za-z_][A-Za-z0-9_]{0,127}$/ : /^[!#$%&'*+.^_`|~0-9A-Za-z-]{1,128}$/
|
||||
if (allKeys.some(key => !validKey.test(key))) throw new Error(`${kind === 'environment' ? '环境变量' : 'Header'}名称无效;敏感变量名只能填名称,不能填密钥值`)
|
||||
for (const [key, value] of Object.entries(values)) {
|
||||
const declared = keys.find(item => identity(item) === identity(key))
|
||||
const sensitive = /api[_-]?key|token|secret|password|authorization|cookie|credential/i.test(key)
|
||||
if (declared || sensitive) {
|
||||
if (!value || value.length > 32768) throw new Error('密钥值必须为 1–32768 个字符')
|
||||
const secretKey = declared ?? key
|
||||
if (!declared) keys.push(key)
|
||||
secrets.push({ kind, key: secretKey, value })
|
||||
delete values[key]
|
||||
} else if (/host|url|endpoint/i.test(key)) checkUrl(value, '环境变量或 Header 地址')
|
||||
}
|
||||
}
|
||||
return { config, secrets }
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
// @vitest-environment happy-dom
|
||||
import { flushPromises, mount } from '@vue/test-utils'
|
||||
import { createPinia } from 'pinia'
|
||||
import { createMemoryHistory, createRouter } from 'vue-router'
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import type { Plugin } from '@/contracts'
|
||||
import * as pluginService from '@/services/pluginService'
|
||||
import PluginMcpPanel from './PluginMcpPanel.vue'
|
||||
import { useWorkspaceStore } from '@/stores/workspace'
|
||||
|
||||
vi.mock('@/services/pluginService', async (loadOriginal) => {
|
||||
const original = await loadOriginal<typeof import('@/services/pluginService')>()
|
||||
return {
|
||||
...original,
|
||||
getPluginHostStatus: vi.fn(),
|
||||
restartPluginHost: vi.fn(),
|
||||
getPluginSettings: vi.fn(),
|
||||
updatePluginSettings: vi.fn(),
|
||||
putPluginSecret: vi.fn(),
|
||||
deletePluginSecret: vi.fn(),
|
||||
listPluginCommands: vi.fn(),
|
||||
executePluginCommand: vi.fn(),
|
||||
}
|
||||
})
|
||||
|
||||
const plugin: Plugin = {
|
||||
plugin_id: 'mcp-demo',
|
||||
name: 'MCP Demo',
|
||||
version: '1.0.0',
|
||||
description: 'demo',
|
||||
status: 'ready',
|
||||
enabled: true,
|
||||
permissions: [],
|
||||
contributions: [
|
||||
{ type: 'settings_section', id: 'mcp-demo.general', name: 'settings' },
|
||||
{ type: 'command', id: 'mcp-demo.run', name: 'run' },
|
||||
],
|
||||
backend_type: 'mcp',
|
||||
transport: 'stdio',
|
||||
}
|
||||
|
||||
async function render() {
|
||||
const pinia = createPinia()
|
||||
const router = createRouter({
|
||||
history: createMemoryHistory(),
|
||||
routes: [{ path: '/', component: { template: '<div />' } }],
|
||||
})
|
||||
await router.push('/')
|
||||
const wrapper = mount(PluginMcpPanel, {
|
||||
props: { plugin },
|
||||
global: { plugins: [pinia, router], stubs: { AppIcon: true } },
|
||||
})
|
||||
const workspaceStore = useWorkspaceStore(pinia)
|
||||
workspaceStore.vaultId = 'default'
|
||||
workspaceStore.hasVault = true
|
||||
return wrapper
|
||||
}
|
||||
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks()
|
||||
vi.mocked(pluginService.getPluginHostStatus).mockResolvedValue({
|
||||
plugin_id: 'mcp-demo', backend_type: 'mcp', transport: 'stdio',
|
||||
status: 'ready', tools_count: 2, server_name: 'demo',
|
||||
})
|
||||
vi.mocked(pluginService.getPluginSettings).mockResolvedValue({
|
||||
plugin_id: 'mcp-demo',
|
||||
schema_version: 1,
|
||||
fields: [
|
||||
{ key: 'limit', label: '数量', description: '', type: 'number', required: true, options: [] },
|
||||
{ key: 'api_key', label: 'API Key', description: '', type: 'secret', required: true, options: [] },
|
||||
],
|
||||
values: { limit: 5 },
|
||||
secrets: { api_key: { configured: false } },
|
||||
})
|
||||
vi.mocked(pluginService.putPluginSecret).mockResolvedValue({
|
||||
plugin_id: 'mcp-demo', key: 'api_key', configured: true,
|
||||
})
|
||||
vi.mocked(pluginService.listPluginCommands).mockResolvedValue([])
|
||||
})
|
||||
|
||||
describe('PluginMcpPanel', () => {
|
||||
it('loads MCP Host status and exposes restart controls', async () => {
|
||||
const wrapper = await render()
|
||||
await flushPromises()
|
||||
expect(pluginService.getPluginHostStatus).toHaveBeenCalledWith('mcp-demo')
|
||||
expect(wrapper.text()).toContain('demo')
|
||||
expect(wrapper.text()).toContain('工具数量')
|
||||
})
|
||||
|
||||
it('builds settings fields from schema and writes secrets separately', async () => {
|
||||
const wrapper = await render()
|
||||
const settingsTab = wrapper.findAll('button').find((button) => button.text() === '设置与密钥')
|
||||
expect(settingsTab).toBeTruthy()
|
||||
await settingsTab!.trigger('click')
|
||||
await flushPromises()
|
||||
expect(wrapper.text()).toContain('数量')
|
||||
expect(wrapper.text()).toContain('API Key')
|
||||
|
||||
await wrapper.get('input[type="password"]').setValue('secret-only-in-request')
|
||||
const secretButton = wrapper.findAll('button').find((button) => button.text() === '安全保存')
|
||||
expect(secretButton).toBeTruthy()
|
||||
await secretButton!.trigger('click')
|
||||
await flushPromises()
|
||||
|
||||
expect(pluginService.putPluginSecret).toHaveBeenCalledWith('mcp-demo', 'api_key', 'secret-only-in-request')
|
||||
expect((wrapper.get('input[type="password"]').element as HTMLInputElement).value).toBe('')
|
||||
expect(wrapper.text()).toContain('已配置')
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,275 @@
|
||||
<script setup lang="ts">
|
||||
import { Key, Refresh, VideoPlay } from '@element-plus/icons-vue'
|
||||
import { computed, ref, watch } from 'vue'
|
||||
import { useRouter } from 'vue-router'
|
||||
import AppIcon from '@/components/common/AppIcon.vue'
|
||||
import type { Plugin, PluginCommand, PluginHostStatus, PluginSettingField, PluginSettingsSchema } from '@/contracts'
|
||||
import * as pluginService from '@/services/pluginService'
|
||||
import { useEditorStore } from '@/stores/editor'
|
||||
import { usePluginStore } from '@/stores/plugin'
|
||||
import { useWorkspaceStore } from '@/stores/workspace'
|
||||
|
||||
const props = defineProps<{ plugin: Plugin }>()
|
||||
const pluginStore = usePluginStore()
|
||||
const editorStore = useEditorStore()
|
||||
const workspaceStore = useWorkspaceStore()
|
||||
const router = useRouter()
|
||||
const activeTab = ref<'host' | 'settings' | 'commands'>('host')
|
||||
const host = ref<PluginHostStatus | null>(null)
|
||||
const schema = ref<PluginSettingsSchema | null>(null)
|
||||
const values = ref<Record<string, unknown>>({})
|
||||
// 明文只停留在组件内存,提交后立即清空。
|
||||
const secrets = ref<Record<string, string>>({})
|
||||
const commands = ref<PluginCommand[]>([])
|
||||
const argumentsByCommand = ref<Record<string, Record<string, unknown>>>({})
|
||||
const loading = ref(false)
|
||||
const busy = ref('')
|
||||
const error = ref('')
|
||||
const notice = ref('')
|
||||
let loadVersion = 0
|
||||
|
||||
const hasSettings = computed(() => props.plugin.contributions.some((item) => item.type === 'settings_section'))
|
||||
const tabs = computed(() => [
|
||||
...(props.plugin.backend_type === 'mcp' ? [{ id: 'host' as const, label: 'MCP Host' }] : []),
|
||||
...(hasSettings.value ? [{ id: 'settings' as const, label: '设置与密钥' }] : []),
|
||||
{ id: 'commands' as const, label: '插件命令' },
|
||||
])
|
||||
|
||||
watch(() => props.plugin.plugin_id, () => {
|
||||
loadVersion++
|
||||
activeTab.value = props.plugin.backend_type === 'mcp' ? 'host' : hasSettings.value ? 'settings' : 'commands'
|
||||
host.value = null
|
||||
schema.value = null
|
||||
values.value = {}
|
||||
secrets.value = {}
|
||||
commands.value = []
|
||||
void loadActive()
|
||||
}, { immediate: true })
|
||||
|
||||
function feedback(message = '') { error.value = message; notice.value = '' }
|
||||
function message(reason: unknown, fallback: string) { return reason instanceof Error ? reason.message : fallback }
|
||||
function formatTime(value?: string | null) { return value ? new Date(value).toLocaleString() : '—' }
|
||||
|
||||
async function selectTab(tab: typeof activeTab.value) {
|
||||
activeTab.value = tab
|
||||
await loadActive()
|
||||
}
|
||||
async function loadActive() {
|
||||
const version = ++loadVersion
|
||||
const pluginId = props.plugin.plugin_id
|
||||
const tab = activeTab.value
|
||||
feedback()
|
||||
loading.value = true
|
||||
try {
|
||||
if (tab === 'host') {
|
||||
const loadedHost = await pluginService.getPluginHostStatus(pluginId)
|
||||
if (version === loadVersion) host.value = loadedHost
|
||||
}
|
||||
if (tab === 'settings') {
|
||||
const loadedSchema = await pluginService.getPluginSettings(pluginId)
|
||||
if (version === loadVersion) {
|
||||
schema.value = loadedSchema
|
||||
values.value = { ...loadedSchema.values }
|
||||
}
|
||||
}
|
||||
if (tab === 'commands') {
|
||||
const loadedCommands = (await pluginService.listPluginCommands()).filter((command) => command.plugin_id === pluginId)
|
||||
if (version === loadVersion) {
|
||||
commands.value = loadedCommands
|
||||
for (const command of loadedCommands) argumentsByCommand.value[command.command_id] = {}
|
||||
}
|
||||
}
|
||||
} catch (reason) {
|
||||
if (version === loadVersion) feedback(message(reason, 'MCP 数据加载失败'))
|
||||
} finally {
|
||||
if (version === loadVersion) loading.value = false
|
||||
}
|
||||
}
|
||||
async function restartHost() {
|
||||
busy.value = 'host'
|
||||
feedback()
|
||||
try {
|
||||
await pluginService.restartPluginHost(props.plugin.plugin_id)
|
||||
host.value = await pluginService.getPluginHostStatus(props.plugin.plugin_id)
|
||||
await pluginStore.loadPlugins()
|
||||
notice.value = 'MCP Host 已重启。'
|
||||
} catch (reason) { feedback(message(reason, 'MCP Host 重启失败')) } finally { busy.value = '' }
|
||||
}
|
||||
function updateValue(field: PluginSettingField, raw: string | boolean) {
|
||||
values.value[field.key] = field.type === 'number' && typeof raw === 'string' ? (raw === '' ? null : Number(raw)) : raw
|
||||
}
|
||||
async function saveSettings() {
|
||||
if (!schema.value) return
|
||||
busy.value = 'settings'
|
||||
feedback()
|
||||
try {
|
||||
schema.value = await pluginService.updatePluginSettings(props.plugin.plugin_id, schema.value.schema_version, values.value)
|
||||
values.value = { ...schema.value.values }
|
||||
notice.value = '普通设置已保存。'
|
||||
} catch (reason) { feedback(message(reason, '设置保存失败')) } finally { busy.value = '' }
|
||||
}
|
||||
async function saveSecret(field: PluginSettingField) {
|
||||
const secret = secrets.value[field.key]?.trim()
|
||||
if (!secret) { feedback('请输入' + field.label); return }
|
||||
busy.value = 'secret:' + field.key
|
||||
feedback()
|
||||
try {
|
||||
const state = await pluginService.putPluginSecret(props.plugin.plugin_id, field.key, secret)
|
||||
if (schema.value) schema.value.secrets[field.key] = { configured: state.configured }
|
||||
secrets.value[field.key] = ''
|
||||
notice.value = field.label + '已加密保存。'
|
||||
} catch (reason) { feedback(message(reason, '密钥保存失败')) } finally { busy.value = '' }
|
||||
}
|
||||
async function deleteSecret(field: PluginSettingField) {
|
||||
if (!confirm('删除已保存的' + field.label + '?')) return
|
||||
busy.value = 'secret:' + field.key
|
||||
feedback()
|
||||
try {
|
||||
const state = await pluginService.deletePluginSecret(props.plugin.plugin_id, field.key)
|
||||
if (schema.value) schema.value.secrets[field.key] = { configured: state.configured }
|
||||
secrets.value[field.key] = ''
|
||||
notice.value = field.label + '已删除。'
|
||||
} catch (reason) { feedback(message(reason, '密钥删除失败')) } finally { busy.value = '' }
|
||||
}
|
||||
function properties(command: PluginCommand): Record<string, Record<string, unknown>> {
|
||||
const result = command.parameters.properties
|
||||
return result && typeof result === 'object' && !Array.isArray(result) ? result as Record<string, Record<string, unknown>> : {}
|
||||
}
|
||||
function required(command: PluginCommand, key: string) {
|
||||
return Array.isArray(command.parameters.required) && command.parameters.required.includes(key)
|
||||
}
|
||||
function commandAvailable(command: PluginCommand) {
|
||||
if (!command.enabled) return false
|
||||
return command.when.every((condition) => {
|
||||
if (condition === 'workspace.has_vault') return Boolean(workspaceStore.vaultId)
|
||||
if (condition === 'editor.has_note') return Boolean(editorStore.currentNoteId)
|
||||
// Plugin 详情页不冒充编辑器选区;选区命令应从命令面板或编辑器挂载点执行。
|
||||
if (condition === 'editor.has_selection') return false
|
||||
return false
|
||||
})
|
||||
}
|
||||
function updateArgument(commandId: string, key: string, raw: string, definition: Record<string, unknown>) {
|
||||
const target = argumentsByCommand.value[commandId] ??= {}
|
||||
if (definition.type === 'number' || definition.type === 'integer') target[key] = raw === '' ? undefined : Number(raw)
|
||||
else if (definition.type === 'boolean') target[key] = raw === 'true'
|
||||
else target[key] = raw
|
||||
}
|
||||
async function execute(command: PluginCommand) {
|
||||
busy.value = command.command_id
|
||||
feedback()
|
||||
try {
|
||||
const result = await pluginService.executePluginCommand(command.command_id, argumentsByCommand.value[command.command_id] ?? {}, {
|
||||
vault_id: workspaceStore.hasVault ? workspaceStore.vaultId : null,
|
||||
note_id: editorStore.currentNoteId,
|
||||
file_path: editorStore.currentFilePath,
|
||||
selection: null,
|
||||
})
|
||||
if (result.effect.type === 'notification') notice.value = result.effect.payload.message
|
||||
else if (result.effect.type === 'job') notice.value = '后台任务已创建:' + result.effect.payload.job_id
|
||||
else if (result.effect.type === 'navigate') {
|
||||
const routes: Record<string, string> = {
|
||||
'vault-entry': '/', workspace: '/workspace', search: '/search', chat: '/chat',
|
||||
agent: '/agent/runs', tasks: '/tasks', skills: '/extensions/skills',
|
||||
plugins: '/extensions/plugins', themes: '/themes', settings: '/settings',
|
||||
}
|
||||
await router.push(routes[result.effect.payload.route])
|
||||
} else if (result.effect.type === 'refresh') {
|
||||
await loadActive()
|
||||
notice.value = '相关数据已刷新。'
|
||||
} else notice.value = '命令执行完成。'
|
||||
} catch (reason) { feedback(message(reason, '命令执行失败')) } finally { busy.value = '' }
|
||||
}
|
||||
</script>
|
||||
|
||||
<template>
|
||||
<section class="mcp-panel">
|
||||
<nav class="mcp-tabs" aria-label="MCP 与 Plugin 配置">
|
||||
<button v-for="tab in tabs" :key="tab.id" :class="{ active: activeTab === tab.id }" @click="selectTab(tab.id)">{{ tab.label }}</button>
|
||||
</nav>
|
||||
<div v-if="error" class="error-banner">{{ error }}</div>
|
||||
<div v-if="notice" class="notice-banner">{{ notice }}</div>
|
||||
|
||||
<div v-if="activeTab === 'host'" class="mcp-section">
|
||||
<div class="section-head"><div><h3>MCP Host 状态</h3><p>查看协议协商、运行状态与 Host 错误。</p></div><div class="inline-actions"><button class="button-secondary" :disabled="loading" @click="loadActive"><AppIcon :icon="Refresh" :size="15" />刷新</button><button class="button-primary" :disabled="busy === 'host' || !plugin.enabled" @click="restartHost">{{ busy === 'host' ? '重启中…' : '重启 Host' }}</button></div></div>
|
||||
<div v-if="host" class="status-grid">
|
||||
<div><span>状态</span><strong><i class="status-dot" :class="host.status"></i>{{ host.status }}</strong></div>
|
||||
<div><span>服务</span><strong>{{ host.server_name || '—' }} {{ host.server_version || '' }}</strong></div>
|
||||
<div><span>协议版本</span><strong>{{ host.protocol_version || '—' }}</strong></div>
|
||||
<div><span>工具数量</span><strong>{{ host.tools_count }}</strong></div>
|
||||
<div><span>启动时间</span><strong>{{ formatTime(host.started_at) }}</strong></div>
|
||||
<div><span>最后心跳</span><strong>{{ formatTime(host.last_seen_at) }}</strong></div>
|
||||
</div>
|
||||
<div v-else-if="loading" class="empty-state">正在读取 Host 状态…</div>
|
||||
<div v-if="host?.error" class="error-banner host-error">{{ host.error }}</div>
|
||||
<p class="security-hint">当前仅运行插件清单声明的 stdio MCP Server,不开放任意 Shell 命令和环境变量编辑。</p>
|
||||
</div>
|
||||
|
||||
<div v-else-if="activeTab === 'settings'" class="mcp-section">
|
||||
<div class="section-head"><div><h3>设置与密钥</h3><p>表单由后端 Schema 生成;密钥不会被读取或回显。</p></div><button class="button-primary" :disabled="!schema || busy === 'settings'" @click="saveSettings">{{ busy === 'settings' ? '保存中…' : '保存普通设置' }}</button></div>
|
||||
<div v-if="schema" class="settings-list">
|
||||
<div v-for="field in schema.fields" :key="field.key" class="setting-row">
|
||||
<div class="field-copy"><label :for="'plugin-setting-' + field.key"><AppIcon v-if="field.type === 'secret'" :icon="Key" :size="15" />{{ field.label }}<em v-if="field.required">必填</em></label><p>{{ field.description || (field.type === 'secret' ? '加密保存,不在页面回显。' : '') }}</p></div>
|
||||
<template v-if="field.type === 'secret'">
|
||||
<div class="secret-control"><input :id="'plugin-setting-' + field.key" :value="secrets[field.key] || ''" class="input" type="password" autocomplete="new-password" :placeholder="schema.secrets[field.key]?.configured ? '已配置;输入新值可替换' : '输入密钥'" @input="secrets[field.key] = ($event.target as HTMLInputElement).value"><button class="button-secondary" :disabled="!secrets[field.key]?.trim() || busy === 'secret:' + field.key" @click="saveSecret(field)">安全保存</button><button v-if="schema.secrets[field.key]?.configured" class="button-danger" @click="deleteSecret(field)">删除</button></div>
|
||||
<span class="secret-state" :class="{ configured: schema.secrets[field.key]?.configured }">{{ schema.secrets[field.key]?.configured ? '已配置' : '未配置' }}</span>
|
||||
</template>
|
||||
<template v-else-if="field.type === 'boolean'"><label class="check-control"><input :id="'plugin-setting-' + field.key" type="checkbox" :checked="Boolean(values[field.key])" @change="updateValue(field, ($event.target as HTMLInputElement).checked)">{{ values[field.key] ? '开启' : '关闭' }}</label></template>
|
||||
<template v-else-if="field.type === 'select'"><select :id="'plugin-setting-' + field.key" class="select" :value="values[field.key]" @change="updateValue(field, ($event.target as HTMLSelectElement).value)"><option v-for="option in field.options" :key="option" :value="option">{{ option }}</option></select></template>
|
||||
<template v-else><input :id="'plugin-setting-' + field.key" class="input" :type="field.type === 'number' ? 'number' : 'text'" :min="field.minimum ?? undefined" :max="field.maximum ?? undefined" :required="field.required" :value="values[field.key] ?? ''" @input="updateValue(field, ($event.target as HTMLInputElement).value)"></template>
|
||||
</div>
|
||||
</div>
|
||||
<div v-else-if="loading" class="empty-state">正在读取 Plugin 设置…</div>
|
||||
</div>
|
||||
|
||||
<div v-else class="mcp-section">
|
||||
<div class="section-head"><div><h3>Plugin 命令</h3><p>执行该 Plugin 注册的受控 Command Contribution。</p></div><button class="button-secondary" :disabled="loading" @click="loadActive"><AppIcon :icon="Refresh" :size="15" />刷新</button></div>
|
||||
<div v-if="commands.length" class="command-list">
|
||||
<article v-for="command in commands" :key="command.command_id" class="item-card command-card">
|
||||
<div class="command-head"><div><strong>{{ command.title }}</strong><p>{{ command.description || command.command_id }}</p></div><span class="badge" :class="{ success: commandAvailable(command), warning: command.enabled && !commandAvailable(command) }">{{ commandAvailable(command) ? '可执行' : command.enabled ? '缺少上下文' : '不可用' }}</span></div>
|
||||
<div v-if="Object.keys(properties(command)).length" class="command-fields">
|
||||
<label v-for="(definition, key) in properties(command)" :key="key" class="field"><span>{{ String(definition.title || key) }}<em v-if="required(command, key)">必填</em></span><select v-if="Array.isArray(definition.enum)" class="select" @change="updateArgument(command.command_id, key, ($event.target as HTMLSelectElement).value, definition)"><option value="">请选择</option><option v-for="option in definition.enum" :key="String(option)" :value="String(option)">{{ option }}</option></select><select v-else-if="definition.type === 'boolean'" class="select" @change="updateArgument(command.command_id, key, ($event.target as HTMLSelectElement).value, definition)"><option value="false">否</option><option value="true">是</option></select><input v-else class="input" :type="definition.type === 'number' || definition.type === 'integer' ? 'number' : 'text'" @input="updateArgument(command.command_id, key, ($event.target as HTMLInputElement).value, definition)"></label>
|
||||
</div>
|
||||
<button class="button-primary command-run" :disabled="!commandAvailable(command) || busy === command.command_id" @click="execute(command)"><AppIcon :icon="VideoPlay" :size="15" />{{ busy === command.command_id ? '执行中…' : '执行命令' }}</button>
|
||||
</article>
|
||||
</div>
|
||||
<div v-else-if="!loading" class="empty-state"><div><strong>没有可用命令</strong><p>启用 Plugin 后,已注册的命令会出现在这里。</p></div></div>
|
||||
</div>
|
||||
</section>
|
||||
</template>
|
||||
|
||||
<style scoped>
|
||||
.mcp-panel { margin-top: var(--space-xl); padding-top: var(--space-xl); border-top: 1px solid var(--color-border-default); }
|
||||
.mcp-tabs { display: flex; gap: var(--space-xs); margin-bottom: var(--space-xl); padding: var(--space-xs); border: 1px solid var(--color-border-default); border-radius: var(--radius-lg); background: var(--color-background-secondary); }
|
||||
.mcp-tabs button { padding: 9px var(--space-md); border-radius: var(--radius-md); color: var(--color-text-secondary); }
|
||||
.mcp-tabs button:hover { background: var(--color-background-hover); }
|
||||
.mcp-tabs button.active { background: var(--color-surface-primary); color: var(--color-accent-primary); box-shadow: var(--shadow-sm); }
|
||||
.mcp-section { min-height: 220px; }
|
||||
.section-head, .command-head { display: flex; align-items: flex-start; justify-content: space-between; gap: var(--space-md); margin-bottom: var(--space-lg); }
|
||||
.section-head p, .command-head p { margin-top: var(--space-xs); color: var(--color-text-tertiary); font-size: var(--font-size-sm); }
|
||||
.section-head button, .command-run { display: inline-flex; align-items: center; gap: var(--space-xs); }
|
||||
.status-grid { display: grid; grid-template-columns: repeat(auto-fit, minmax(165px, 1fr)); gap: var(--space-sm); }
|
||||
.status-grid > div { display: grid; gap: var(--space-xs); padding: var(--space-md); border: 1px solid var(--color-border-subtle); border-radius: var(--radius-md); background: var(--color-background-secondary); }
|
||||
.status-grid span { color: var(--color-text-tertiary); font-size: var(--font-size-xs); }
|
||||
.status-grid strong { display: flex; align-items: center; gap: var(--space-xs); font-size: var(--font-size-sm); }
|
||||
.status-dot { width: 8px; height: 8px; border-radius: 50%; background: var(--color-text-tertiary); }
|
||||
.status-dot.ready { background: var(--color-success); box-shadow: 0 0 0 4px var(--color-success-soft); }
|
||||
.status-dot.error, .status-dot.unhealthy { background: var(--color-error); box-shadow: 0 0 0 4px var(--color-error-soft); }
|
||||
.status-dot.starting { background: var(--color-warning); box-shadow: 0 0 0 4px var(--color-warning-soft); }
|
||||
.security-hint { margin-top: var(--space-lg); padding: var(--space-md); border-left: 3px solid var(--color-info); background: var(--color-info-soft); color: var(--color-text-secondary); font-size: var(--font-size-sm); }
|
||||
.host-error { margin-top: var(--space-lg); }
|
||||
.settings-list { display: grid; }
|
||||
.setting-row { display: grid; grid-template-columns: minmax(180px, .9fr) minmax(260px, 1.1fr) auto; align-items: center; gap: var(--space-lg); padding: var(--space-lg) 0; border-bottom: 1px solid var(--color-border-subtle); }
|
||||
.field-copy label { display: flex; align-items: center; gap: var(--space-xs); font-weight: 650; }
|
||||
.field-copy p { margin-top: var(--space-xs); color: var(--color-text-tertiary); font-size: var(--font-size-sm); }
|
||||
em { margin-left: var(--space-xs); color: var(--color-error); font-size: var(--font-size-xs); font-style: normal; }
|
||||
.secret-control { display: flex; gap: var(--space-xs); }
|
||||
.secret-state { color: var(--color-text-tertiary); font-size: var(--font-size-xs); }
|
||||
.secret-state.configured { color: var(--color-success); }
|
||||
.check-control { display: flex; align-items: center; gap: var(--space-sm); color: var(--color-text-secondary); }
|
||||
.check-control input { width: 18px; height: 18px; accent-color: var(--color-accent-primary); }
|
||||
.command-list, .command-card { display: grid; gap: var(--space-sm); }
|
||||
.command-card:hover { transform: none; }
|
||||
.command-fields { display: grid; grid-template-columns: repeat(auto-fit, minmax(180px, 1fr)); gap: var(--space-md); }
|
||||
.command-run { justify-self: end; }
|
||||
@media (max-width: 800px) { .mcp-tabs { overflow-x: auto; } .mcp-tabs button { flex: 0 0 auto; } .setting-row { grid-template-columns: 1fr; gap: var(--space-sm); } .secret-control { flex-wrap: wrap; } }
|
||||
</style>
|
||||
@@ -1,6 +1,7 @@
|
||||
<script setup lang="ts">
|
||||
import { Connection } from '@element-plus/icons-vue'
|
||||
import AppIcon from '@/components/common/AppIcon.vue'
|
||||
import PluginMcpPanel from './PluginMcpPanel.vue'
|
||||
import { onMounted, ref } from 'vue'
|
||||
import { usePluginStore } from '@/stores/plugin'
|
||||
|
||||
@@ -16,7 +17,7 @@ async function uninstall(id: string, name: string) { if (!confirm(`卸载“${na
|
||||
|
||||
<template>
|
||||
<section class="feature-page">
|
||||
<header class="feature-header"><div><h1>Plugin 管理</h1><p>管理插件生命周期、权限和受控 Contribution。</p></div><button class="button-primary" @click="install">安装 Plugin</button></header>
|
||||
<header class="feature-header"><div><h1>Plugin 与 MCP</h1><p>管理插件生命周期、MCP Host、权限和受控 Contribution。</p></div><button class="button-primary" @click="install">安装 Plugin</button></header>
|
||||
<div v-if="pluginStore.error || actionError" class="error-banner">{{ pluginStore.error || actionError }}</div>
|
||||
<div v-if="pluginStore.selectedPlugin" class="panel">
|
||||
<div class="detail-head"><div><span class="badge" :class="{ success: pluginStore.selectedPlugin.status === 'ready', error: pluginStore.selectedPlugin.status === 'error', warning: pluginStore.selectedPlugin.status === 'permission_required' }">{{ pluginStore.selectedPlugin.status }}</span><h2>{{ pluginStore.selectedPlugin.icon }} {{ pluginStore.selectedPlugin.name }}</h2><p class="muted">v{{ pluginStore.selectedPlugin.version }} · {{ pluginStore.selectedPlugin.backend_type || 'none' }}/{{ pluginStore.selectedPlugin.transport || 'none' }}</p></div><div class="inline-actions"><button v-if="pluginStore.selectedPlugin.status === 'permission_required'" class="button-primary" @click="grant(pluginStore.selectedPlugin.plugin_id, pluginStore.selectedPlugin.permissions)">授权权限</button><button class="button-secondary" @click="toggle(pluginStore.selectedPlugin.plugin_id, pluginStore.selectedPlugin.enabled)">{{ pluginStore.selectedPlugin.enabled ? '停用' : '启用' }}</button><button class="button-danger" @click="uninstall(pluginStore.selectedPlugin.plugin_id, pluginStore.selectedPlugin.name)">卸载</button></div></div>
|
||||
@@ -24,6 +25,7 @@ async function uninstall(id: string, name: string) { if (!confirm(`卸载“${na
|
||||
<div class="detail-grid"><div><h3>权限</h3><div class="tag-list"><span v-for="permission in pluginStore.selectedPlugin.permissions" :key="permission" class="badge warning">{{ permission }}</span></div></div><div><h3>Contribution</h3><div class="contribution-list"><div v-for="item in pluginStore.selectedPlugin.contributions" :key="item.id" class="item-card"><span class="badge info">{{ item.type }}</span><strong>{{ item.name }}</strong><p class="subtle">{{ item.description || item.id }}</p></div></div></div></div>
|
||||
<div v-if="pluginStore.selectedPlugin.last_error" class="error-banner last-error">{{ pluginStore.selectedPlugin.last_error }}</div>
|
||||
<div v-if="pluginStore.selectedPlugin.dependent_skills?.length" class="notice-banner last-error">依赖此插件的 Skill:{{ pluginStore.selectedPlugin.dependent_skills.join('、') }}</div>
|
||||
<PluginMcpPanel :plugin="pluginStore.selectedPlugin" />
|
||||
</div>
|
||||
<div v-else class="feature-grid"><article v-for="plugin in pluginStore.plugins" :key="plugin.plugin_id" class="item-card extension-card" @click="pluginStore.selectPlugin(plugin.plugin_id)"><div class="extension-title"><AppIcon :icon="Connection" :size="22" /><div><strong>{{ plugin.name }}</strong><p>v{{ plugin.version }}</p></div><span class="badge" :class="{ success: plugin.status === 'ready', error: plugin.status === 'error', warning: plugin.status === 'permission_required' }">{{ plugin.status }}</span></div><p class="muted">{{ plugin.description }}</p><p class="subtle">{{ plugin.permissions.length }} 项权限 · {{ plugin.contributions.length }} 项 Contribution</p></article></div>
|
||||
</section>
|
||||
|
||||
@@ -0,0 +1,170 @@
|
||||
// @vitest-environment happy-dom
|
||||
import { flushPromises, mount } from '@vue/test-utils'
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import type { ModelRoutingResponse, ProviderConfig } from '@/contracts'
|
||||
import { ApiErrorClass } from '@/services/apiClient'
|
||||
import * as service from '@/services/modelRoutingService'
|
||||
import { listProviders } from '@/services/providerService'
|
||||
import ModelRoutingSettings from './ModelRoutingSettings.vue'
|
||||
|
||||
vi.mock('@/services/modelRoutingService', () => ({ getModelRouting: vi.fn(), saveModelRouting: vi.fn() }))
|
||||
vi.mock('@/services/providerService', () => ({ listProviders: vi.fn() }))
|
||||
const providers: ProviderConfig[] = [
|
||||
{ provider_id: 'p1', provider_type: 'openai_compatible', name: 'Custom API', enabled: true, default_model: 'chat-model', capabilities: {}, has_credential: true },
|
||||
{ provider_id: 'p2', provider_type: 'openai_chat', name: 'OpenAI', enabled: true, default_model: '', capabilities: {}, has_credential: true },
|
||||
{ provider_id: 'responses', provider_type: 'openai_responses', name: 'Responses', enabled: true, default_model: '', capabilities: {}, has_credential: true },
|
||||
{ provider_id: 'anthropic', provider_type: 'anthropic_messages', name: 'Anthropic', enabled: true, default_model: '', capabilities: {}, has_credential: true },
|
||||
{ provider_id: 'ollama', provider_type: 'ollama', name: 'Ollama', enabled: true, default_model: '', capabilities: {}, has_credential: false },
|
||||
{ provider_id: 'disabled', provider_type: 'openai_chat', name: 'Disabled', enabled: false, default_model: '', capabilities: {}, has_credential: true },
|
||||
]
|
||||
const initial: ModelRoutingResponse = { config: { version: 3, embedding: null, transcription: null, speaker_matching: null }, local_backends: [
|
||||
{ capability: 'embedding', status: 'placeholder', message: 'hash fallback' },
|
||||
{ capability: 'transcription', status: 'not_installed', message: 'ASR not installed' },
|
||||
{ capability: 'speaker_matching', status: 'not_installed', message: 'speaker not installed' },
|
||||
] }
|
||||
const wrappers: ReturnType<typeof mount>[] = []
|
||||
async function render() {
|
||||
const wrapper = mount(ModelRoutingSettings)
|
||||
wrappers.push(wrapper)
|
||||
await flushPromises()
|
||||
return wrapper
|
||||
}
|
||||
beforeEach(() => {
|
||||
vi.resetAllMocks()
|
||||
vi.mocked(service.getModelRouting).mockResolvedValue(structuredClone(initial))
|
||||
vi.mocked(listProviders).mockResolvedValue(providers)
|
||||
vi.mocked(service.saveModelRouting).mockImplementation(async config => ({ ...initial, config: { ...config, version: config.version + 1 } }))
|
||||
})
|
||||
afterEach(() => { wrappers.splice(0).forEach(wrapper => wrapper.unmount()) })
|
||||
|
||||
describe('ModelRoutingSettings', () => {
|
||||
it('loads local selections honestly, explains index rebuilds, and disables incompatible providers', async () => {
|
||||
const wrapper = await render()
|
||||
expect(wrapper.findAll('select').map(select => (select.element as HTMLSelectElement).value)).toEqual(['', '', ''])
|
||||
expect(wrapper.text()).toContain('当前为占位实现')
|
||||
expect(wrapper.text()).toContain('真实本地 ASR 尚未接入')
|
||||
expect(wrapper.text()).toContain('真实本地说话人匹配尚未接入')
|
||||
expect(wrapper.text()).toContain('重建全部')
|
||||
expect(wrapper.text()).toContain('重建完成前继续使用本地检索')
|
||||
expect(wrapper.text()).toContain('不是 OpenAI 标准接口')
|
||||
for (const id of ['responses', 'anthropic', 'ollama', 'disabled']) expect(wrapper.get(`option[value="${id}"]`).attributes()).toHaveProperty('disabled')
|
||||
expect(wrapper.get('option[value="p1"]').attributes()).not.toHaveProperty('disabled')
|
||||
})
|
||||
|
||||
it('saves all three independent bindings with expected version and remote dimensions', async () => {
|
||||
const wrapper = await render()
|
||||
for (const capability of ['embedding', 'transcription', 'speaker_matching']) {
|
||||
const card = wrapper.get(`[data-capability="${capability}"]`)
|
||||
await card.get('select').setValue('p1')
|
||||
await card.get('[data-field="model"]').setValue(`${capability}-model`)
|
||||
}
|
||||
await wrapper.get('[data-field="dimensions"]').setValue('3072')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(service.saveModelRouting).toHaveBeenCalledWith({ version: 3,
|
||||
embedding: { provider_id: 'p1', model: 'embedding-model', endpoint: '/embeddings', dimensions: 3072 },
|
||||
transcription: { provider_id: 'p1', model: 'transcription-model', endpoint: '/audio/transcriptions' },
|
||||
speaker_matching: { provider_id: 'p1', model: 'speaker_matching-model', endpoint: '/audio/speaker-matches' },
|
||||
})
|
||||
expect(wrapper.text()).toContain('配置版本 4')
|
||||
expect(wrapper.text()).toContain('模型路由已保存')
|
||||
await wrapper.get('[data-capability="embedding"] select').setValue('')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(service.saveModelRouting).toHaveBeenLastCalledWith(expect.objectContaining({ version: 4, embedding: null }))
|
||||
})
|
||||
|
||||
it('supports omitted dimensions and clears stale models/endpoints when switching providers', async () => {
|
||||
const wrapper = await render()
|
||||
const card = wrapper.get('[data-capability="embedding"]')
|
||||
await card.get('select').setValue('p1')
|
||||
await card.get('[data-field="model"]').setValue('custom-embedding')
|
||||
await card.get('[data-field="endpoint"]').setValue('/custom/embeddings')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(service.saveModelRouting).toHaveBeenCalledWith(expect.objectContaining({ embedding: { provider_id: 'p1', model: 'custom-embedding', endpoint: '/custom/embeddings', dimensions: null } }))
|
||||
await card.get('select').setValue('p2')
|
||||
expect((card.get('[data-field="model"]').element as HTMLInputElement).value).toBe('')
|
||||
expect((card.get('[data-field="endpoint"]').element as HTMLInputElement).value).toBe('/embeddings')
|
||||
})
|
||||
|
||||
it('shows a loading state and does not offer a default local configuration after load failure', async () => {
|
||||
let fail!: (error: Error) => void
|
||||
vi.mocked(service.getModelRouting).mockReturnValueOnce(new Promise((_, reject) => { fail = reject }))
|
||||
const wrapper = await render()
|
||||
expect(wrapper.text()).toContain('正在加载模型路由')
|
||||
expect(wrapper.find('form').exists()).toBe(false)
|
||||
fail(new Error('offline'))
|
||||
await flushPromises()
|
||||
expect(wrapper.text()).toContain('加载失败:offline')
|
||||
expect(wrapper.find('form').exists()).toBe(false)
|
||||
await wrapper.get('button').trigger('click')
|
||||
await flushPromises()
|
||||
expect(wrapper.find('form').exists()).toBe(true)
|
||||
})
|
||||
|
||||
it('keeps unsaved input on save failure and retries without inventing a new version', async () => {
|
||||
const wrapper = await render()
|
||||
const card = wrapper.get('[data-capability="transcription"]')
|
||||
await card.get('select').setValue('p1')
|
||||
await card.get('[data-field="model"]').setValue('asr-model')
|
||||
vi.mocked(service.saveModelRouting).mockRejectedValueOnce(new Error('disk full'))
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(wrapper.text()).toContain('保存失败:disk full')
|
||||
expect((card.get('[data-field="model"]').element as HTMLInputElement).value).toBe('asr-model')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(service.saveModelRouting).toHaveBeenLastCalledWith(expect.objectContaining({ version: 3 }))
|
||||
})
|
||||
|
||||
it('blocks overwrite after a conflict until explicitly reloading the latest configuration', async () => {
|
||||
const wrapper = await render()
|
||||
vi.mocked(service.saveModelRouting).mockRejectedValueOnce(new ApiErrorClass('MODEL_ROUTING_VERSION_CONFLICT', 'stale'))
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(wrapper.text()).toContain('配置版本冲突')
|
||||
expect(wrapper.get('button[type="submit"]').attributes()).toHaveProperty('disabled')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
expect(service.saveModelRouting).toHaveBeenCalledTimes(1)
|
||||
vi.mocked(service.getModelRouting).mockResolvedValueOnce({ ...initial, config: { ...initial.config, version: 8 } })
|
||||
await wrapper.findAll('button').find(button => button.text().includes('放弃当前输入'))!.trigger('click')
|
||||
await flushPromises()
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(service.saveModelRouting).toHaveBeenLastCalledWith(expect.objectContaining({ version: 8 }))
|
||||
})
|
||||
|
||||
it('prevents invalid dimensions, endpoints, and missing provider bindings from being saved', async () => {
|
||||
vi.mocked(service.getModelRouting).mockResolvedValueOnce({ ...initial, config: { ...initial.config, embedding: { provider_id: 'missing', model: 'old-model', endpoint: '/embeddings' } } })
|
||||
const wrapper = await render()
|
||||
expect(wrapper.text()).toContain('原提供商已不可用')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
expect(service.saveModelRouting).not.toHaveBeenCalled()
|
||||
const card = wrapper.get('[data-capability="embedding"]')
|
||||
await card.get('select').setValue('p1')
|
||||
await card.get('[data-field="model"]').setValue('embedding-model')
|
||||
await card.get('[data-field="dimensions"]').setValue('1.5')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
expect(wrapper.text()).toContain('嵌入维度必须为 1–16384 的整数')
|
||||
await card.get('[data-field="dimensions"]').setValue('16385')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
expect(service.saveModelRouting).not.toHaveBeenCalled()
|
||||
expect(card.get('[data-field="dimensions"]').attributes('max')).toBe('16384')
|
||||
await card.get('[data-field="dimensions"]').setValue('')
|
||||
await card.get('[data-field="endpoint"]').setValue('https://example.test/embeddings')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
expect(wrapper.text()).toContain('Endpoint 必须是以 / 开头的相对路径')
|
||||
expect(service.saveModelRouting).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it('labels injected ready local backends accurately', async () => {
|
||||
vi.mocked(service.getModelRouting).mockResolvedValueOnce({ ...initial, local_backends: [{ capability: 'transcription', status: 'ready', message: 'Local ASR ready' }] })
|
||||
const wrapper = await render()
|
||||
const card = wrapper.get('[data-capability="transcription"]')
|
||||
expect(card.get('option[value=""]').text()).toBe('本地 · 已就绪')
|
||||
expect(card.text()).toContain('本地后端已就绪')
|
||||
expect(card.text()).toContain('Local ASR ready')
|
||||
expect(card.text()).not.toContain('真实本地 ASR 尚未接入')
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,160 @@
|
||||
<script setup lang="ts">
|
||||
import { computed, onBeforeUnmount, onMounted, reactive, ref } from 'vue'
|
||||
import type { ModelBinding, ModelRoutingConfig, ModelRoutingResponse, ProviderConfig, RoutingCapability } from '@/contracts'
|
||||
import { getModelRouting, saveModelRouting } from '@/services/modelRoutingService'
|
||||
import { listProviders } from '@/services/providerService'
|
||||
import { ApiErrorClass } from '@/services/apiClient'
|
||||
|
||||
const capabilities: Array<{ id: RoutingCapability; name: string; endpoint: string; placeholder: string; local: string }> = [
|
||||
{ id: 'embedding', name: '向量嵌入 · Embedding', endpoint: '/embeddings', placeholder: '例如 text-embedding-3-small', local: '当前为占位实现,尚未接入真实本地嵌入模型。' },
|
||||
{ id: 'transcription', name: '语音转写 · Transcription', endpoint: '/audio/transcriptions', placeholder: '输入转写模型 ID', local: '真实本地 ASR 尚未接入,等待阶段 F;当前无法进行本地语音识别。' },
|
||||
{ id: 'speaker_matching', name: '说话人匹配 · Speaker matching', endpoint: '/audio/speaker-matches', placeholder: '输入说话人匹配模型 ID', local: '真实本地说话人匹配尚未接入,等待阶段 F;当前无法进行本地声纹匹配。' },
|
||||
]
|
||||
type Draft = { provider_id: string; model: string; endpoint: string; dimensions: string | number }
|
||||
const drafts = reactive(Object.fromEntries(capabilities.map(item => [item.id, { provider_id: '', model: '', endpoint: item.endpoint, dimensions: '' }])) as Record<RoutingCapability, Draft>)
|
||||
const providers = ref<ProviderConfig[]>([])
|
||||
const response = ref<ModelRoutingResponse | null>(null)
|
||||
const loading = ref(false)
|
||||
const saving = ref(false)
|
||||
const error = ref('')
|
||||
const saved = ref(false)
|
||||
const conflict = ref(false)
|
||||
let active = true
|
||||
const eligible = (provider: ProviderConfig) => provider.enabled && ['openai_chat', 'openai_compatible'].includes(provider.provider_type)
|
||||
const available = computed(() => providers.value.filter(eligible))
|
||||
const unavailable = computed(() => providers.value.filter(provider => !eligible(provider)))
|
||||
const localBackend = (capability: RoutingCapability) => response.value?.local_backends.find(item => item.capability === capability)
|
||||
const localLabel = (capability: RoutingCapability) => {
|
||||
const status = localBackend(capability)?.status
|
||||
return status === 'ready' ? '已就绪' : status === 'placeholder' ? '占位实现' : '尚未接入'
|
||||
}
|
||||
const protocols = [
|
||||
{ id: 'openai_chat', label: 'OpenAI Chat' }, { id: 'openai_compatible', label: 'OpenAI Compatible' },
|
||||
{ id: 'openai_responses', label: 'Responses' }, { id: 'anthropic_messages', label: 'Anthropic' }, { id: 'ollama', label: 'Ollama' },
|
||||
]
|
||||
|
||||
function applyResponse(result: ModelRoutingResponse) {
|
||||
response.value = result
|
||||
for (const item of capabilities) {
|
||||
const binding = result.config[item.id]
|
||||
Object.assign(drafts[item.id], { provider_id: binding?.provider_id ?? '', model: binding?.model ?? '', endpoint: binding?.endpoint ?? item.endpoint, dimensions: binding?.dimensions?.toString() ?? '' })
|
||||
}
|
||||
}
|
||||
|
||||
async function load() {
|
||||
if (loading.value || saving.value) return
|
||||
loading.value = true
|
||||
error.value = ''
|
||||
saved.value = false
|
||||
try {
|
||||
const [routing, items] = await Promise.all([getModelRouting(), listProviders()])
|
||||
if (!active) return
|
||||
providers.value = items
|
||||
applyResponse(routing)
|
||||
conflict.value = false
|
||||
} catch (reason) {
|
||||
if (active) error.value = `加载失败:${reason instanceof Error ? reason.message : '无法读取模型路由或提供商'}`
|
||||
} finally { loading.value = false }
|
||||
}
|
||||
|
||||
onMounted(load)
|
||||
onBeforeUnmount(() => { active = false })
|
||||
|
||||
function changeProvider(capability: RoutingCapability) {
|
||||
const draft = drafts[capability]
|
||||
draft.model = ''
|
||||
draft.dimensions = ''
|
||||
draft.endpoint = capabilities.find(item => item.id === capability)!.endpoint
|
||||
saved.value = false
|
||||
}
|
||||
|
||||
function bindingFor(capability: RoutingCapability): ModelBinding | null {
|
||||
const draft = drafts[capability]
|
||||
if (!draft.provider_id) return null
|
||||
if (!available.value.some(provider => provider.provider_id === draft.provider_id)) throw new Error('请选择已启用且协议可用的提供商,或切换到本地。')
|
||||
if (!draft.model.trim()) throw new Error('请填写所选 API 的模型 ID。')
|
||||
if (!/^\/[A-Za-z0-9_/-]+$/.test(draft.endpoint) || draft.endpoint.startsWith('//')) throw new Error('Endpoint 必须是以 / 开头的相对路径,只能包含字母、数字、下划线、连字符和 /。')
|
||||
const binding: ModelBinding = { provider_id: draft.provider_id, model: draft.model.trim(), endpoint: draft.endpoint }
|
||||
if (capability === 'embedding') {
|
||||
const dimension = String(draft.dimensions).trim()
|
||||
if (dimension && (!/^\d+$/.test(dimension) || !Number.isSafeInteger(Number(dimension)) || Number(dimension) < 1 || Number(dimension) > 16384)) throw new Error('嵌入维度必须为 1–16384 的整数,或留空使用 API 默认值。')
|
||||
binding.dimensions = dimension ? Number(dimension) : null
|
||||
}
|
||||
return binding
|
||||
}
|
||||
|
||||
async function save() {
|
||||
if (!response.value || loading.value || saving.value || conflict.value) return
|
||||
saving.value = true
|
||||
error.value = ''
|
||||
saved.value = false
|
||||
try {
|
||||
const config: ModelRoutingConfig = {
|
||||
version: response.value.config.version,
|
||||
embedding: bindingFor('embedding'), transcription: bindingFor('transcription'), speaker_matching: bindingFor('speaker_matching'),
|
||||
}
|
||||
const result = await saveModelRouting(config)
|
||||
if (active) { applyResponse(result); saved.value = true }
|
||||
} catch (reason) {
|
||||
if (!active) return
|
||||
conflict.value = reason instanceof ApiErrorClass && /CONFLICT|VERSION|HTTP_409/i.test(reason.code)
|
||||
error.value = conflict.value
|
||||
? '配置版本冲突:其他窗口已修改路由。当前输入尚未保存,请重新加载最新配置后再编辑。'
|
||||
: `保存失败:${reason instanceof Error ? reason.message : '请重试'}`
|
||||
} finally { saving.value = false }
|
||||
}
|
||||
</script>
|
||||
|
||||
<template>
|
||||
<section class="routing-settings" aria-labelledby="routing-title" :aria-busy="loading || saving">
|
||||
<div><h2 id="routing-title">能力模型路由</h2><p class="subtle">向量嵌入、语音转写和说话人匹配分别选择提供商与模型,独立于默认聊天模型。API Key 在「模型提供商」中管理。</p></div>
|
||||
<p class="subtle">未选择提供商即使用本地路径。API 请求失败、配置不可用或响应无效时,服务端会回退到当前本地处理;本地占位不代表真实模型已接入。</p>
|
||||
<p v-if="loading" role="status">正在加载模型路由…</p>
|
||||
<div v-if="error" class="error-banner" role="alert">{{ error }}</div>
|
||||
<div class="inline-actions"><button type="button" class="button-secondary" :disabled="loading || saving" @click="load">{{ conflict ? '放弃当前输入并加载最新配置' : response ? '重新加载(放弃未保存更改)' : '重试加载' }}</button><span v-if="response" class="subtle">配置版本 {{ response.config.version }}</span></div>
|
||||
<form v-if="response" @submit.prevent="save" @input="saved = false" @change="saved = false">
|
||||
<fieldset :disabled="loading || saving || conflict">
|
||||
<article v-for="capability in capabilities" :key="capability.id" class="routing-card" :data-capability="capability.id">
|
||||
<h3>{{ capability.name }}</h3>
|
||||
<p v-if="capability.id === 'embedding'" class="embedding-notice">更换模型或接口后,请重建全部索引。重建完成前继续使用本地检索。</p>
|
||||
<div class="protocols" aria-label="协议可用性">
|
||||
<span v-for="protocol in protocols" :key="protocol.id" class="badge" :class="{ 'protocol-unavailable': !['openai_chat', 'openai_compatible'].includes(protocol.id) }">{{ protocol.label }}{{ ['openai_chat', 'openai_compatible'].includes(protocol.id) ? ' · 可用' : ' · 不可用' }}</span>
|
||||
</div>
|
||||
<label class="field"><span>处理方式 / 提供商</span><select v-model="drafts[capability.id].provider_id" class="select" data-field="provider" @change="changeProvider(capability.id)">
|
||||
<option value="">本地 · {{ localLabel(capability.id) }}</option>
|
||||
<option v-for="provider in available" :key="provider.provider_id" :value="provider.provider_id">{{ provider.name }} · {{ provider.provider_type }}</option>
|
||||
<option v-for="provider in unavailable" :key="provider.provider_id" :value="provider.provider_id" disabled>{{ provider.name }} · {{ provider.enabled ? '协议不可用' : '未启用' }}</option>
|
||||
<option v-if="drafts[capability.id].provider_id && !providers.some(provider => provider.provider_id === drafts[capability.id].provider_id)" :value="drafts[capability.id].provider_id" disabled>原提供商已不可用 · {{ drafts[capability.id].provider_id }}</option>
|
||||
</select></label>
|
||||
<div v-if="drafts[capability.id].provider_id" class="routing-fields">
|
||||
<label class="field"><span>模型 ID</span><input v-model="drafts[capability.id].model" class="input" data-field="model" :placeholder="capability.placeholder" maxlength="256" required /></label>
|
||||
<label class="field"><span>Endpoint(相对 Base URL)</span><input v-model="drafts[capability.id].endpoint" class="input" data-field="endpoint" :placeholder="capability.endpoint" maxlength="256" required /></label>
|
||||
<label v-if="capability.id === 'embedding'" class="field"><span>向量维度(可选)</span><input v-model="drafts.embedding.dimensions" class="input" data-field="dimensions" type="number" min="1" max="16384" step="1" placeholder="留空使用 API 默认维度" /><small class="subtle">填写模型支持的 1–16384 整数维度,或留空使用 API 默认值。</small></label>
|
||||
</div>
|
||||
<p v-if="capability.id === 'speaker_matching'" class="subtle">说话人匹配使用本应用自定义 HTTP multipart 契约。该端点不是 OpenAI 标准接口;服务需实现对应的说话人匹配请求和响应。</p>
|
||||
<div class="local-status" :class="{ selected: !drafts[capability.id].provider_id }">
|
||||
<strong>{{ drafts[capability.id].provider_id ? '本地回退状态' : '当前本地状态' }}</strong>
|
||||
<p>{{ localBackend(capability.id)?.status === 'ready' ? '本地后端已就绪。' : capability.local }}</p>
|
||||
<p v-for="backend in response.local_backends.filter(item => item.capability === capability.id)" :key="backend.capability" class="subtle"><span class="badge">{{ backend.status === 'ready' ? '已就绪' : backend.status === 'placeholder' ? '占位实现' : '未安装 / 未接入' }}</span> {{ backend.message }}</p>
|
||||
</div>
|
||||
</article>
|
||||
</fieldset>
|
||||
<div class="inline-actions"><button type="submit" class="button-primary" :disabled="loading || saving || conflict">{{ saving ? '保存中…' : '保存模型路由' }}</button><span v-if="saved" role="status">模型路由已保存。</span></div>
|
||||
</form>
|
||||
</section>
|
||||
</template>
|
||||
|
||||
<style scoped>
|
||||
.routing-settings, form, fieldset { display: grid; gap: var(--space-lg); }
|
||||
.routing-settings { border-top: 1px solid var(--color-border-default); padding-top: var(--space-xl); margin-top: var(--space-md); }
|
||||
fieldset { min-width: 0; padding: 0; margin: 0; border: 0; }
|
||||
.routing-card { display: grid; gap: var(--space-md); padding: var(--space-lg); border: 1px solid var(--color-border-default); border-radius: var(--radius-lg); background: var(--color-surface-primary); }
|
||||
.routing-card h3 { margin: 0; }
|
||||
.protocols { display: flex; flex-wrap: wrap; gap: var(--space-xs); }
|
||||
.protocol-unavailable { opacity: .65; }
|
||||
.embedding-notice { padding: var(--space-md); border-radius: var(--radius-md); background: var(--color-accent-soft); color: var(--color-text-primary); }
|
||||
.routing-fields { display: grid; grid-template-columns: repeat(2, minmax(0, 1fr)); gap: var(--space-md); }
|
||||
.local-status { display: grid; gap: var(--space-xs); padding: var(--space-md); background: var(--color-background-secondary); border-radius: var(--radius-md); }
|
||||
.local-status.selected { border-left: 3px solid var(--color-accent-primary); }
|
||||
@media (max-width: 700px) { .routing-fields { grid-template-columns: 1fr; } }
|
||||
</style>
|
||||
@@ -0,0 +1,157 @@
|
||||
// @vitest-environment happy-dom
|
||||
import { flushPromises, mount } from '@vue/test-utils'
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import type { ProviderConfig, ProviderPreset } from '@/contracts'
|
||||
import * as service from '@/services/providerService'
|
||||
import ProviderForm from './ProviderForm.vue'
|
||||
import ProviderPresetSelector from './ProviderPresetSelector.vue'
|
||||
|
||||
vi.mock('@/services/providerService', () => ({ listProviderPresets: vi.fn(), getCredentialStatus: vi.fn(), putCredential: vi.fn(), createProvider: vi.fn(), updateProvider: vi.fn() }))
|
||||
const presets: ProviderPreset[] = [
|
||||
{ preset_id: 'deepseek', name: 'DeepSeek', provider_type: 'openai_compatible', base_url: 'https://deepseek.example.test', default_credential_id: 'shared-deepseek', requires_credential: true, logo_id: 'deepseek' },
|
||||
{ preset_id: 'qwen', name: '通义千问', provider_type: 'openai_compatible', base_url: 'https://qwen.example.test', default_credential_id: 'shared-qwen', requires_credential: true, logo_id: 'qwen' },
|
||||
]
|
||||
const existing: ProviderConfig = { provider_id: 'p1', provider_type: 'openai_compatible', name: 'DeepSeek', base_url: presets[0].base_url, default_model: 'old-model', enabled: true, credential_id: 'old-shared-key', has_credential: true, capabilities: {} }
|
||||
const wrappers: ReturnType<typeof mount>[] = []
|
||||
async function render(provider?: ProviderConfig) {
|
||||
const wrapper = mount(ProviderForm, { props: { provider, models: [{ model_id: 'old-model', name: 'Old', capabilities: {} }] } })
|
||||
wrappers.push(wrapper)
|
||||
await flushPromises()
|
||||
return wrapper
|
||||
}
|
||||
beforeEach(() => {
|
||||
vi.resetAllMocks()
|
||||
vi.mocked(service.listProviderPresets).mockResolvedValue(presets)
|
||||
vi.mocked(service.getCredentialStatus).mockResolvedValue(true)
|
||||
vi.mocked(service.putCredential).mockResolvedValue()
|
||||
vi.mocked(service.createProvider).mockResolvedValue(existing)
|
||||
vi.mocked(service.updateProvider).mockResolvedValue(existing)
|
||||
})
|
||||
afterEach(() => { wrappers.splice(0).forEach(wrapper => wrapper.unmount()) })
|
||||
|
||||
describe('ProviderForm', () => {
|
||||
it('filters compact preset chips and resolves bundled logos', async () => {
|
||||
const wrapper = await render()
|
||||
await wrapper.get('#provider-search').setValue('通义')
|
||||
expect(wrapper.find('[data-preset="deepseek"]').exists()).toBe(false)
|
||||
expect(wrapper.find('[data-preset="qwen"]').exists()).toBe(true)
|
||||
expect(wrapper.get('[data-preset="qwen"] img').attributes('src')).not.toMatch(/^https?:/)
|
||||
})
|
||||
|
||||
it('clears the secret, old model and credential on preset and custom selection', async () => {
|
||||
const wrapper = await render(existing)
|
||||
await wrapper.get('input[type="password"]').setValue('draft-secret')
|
||||
await wrapper.get('[data-preset="qwen"]').trigger('click')
|
||||
expect((wrapper.get('input[type="password"]').element as HTMLInputElement).value).toBe('')
|
||||
expect((wrapper.get('[data-field="model"]').element as HTMLInputElement).value).toBe('')
|
||||
expect(wrapper.findAll('datalist option')).toHaveLength(0)
|
||||
await wrapper.get('form').trigger('submit')
|
||||
expect(service.updateProvider).not.toHaveBeenCalled()
|
||||
expect(wrapper.text()).toContain('请输入 API Key')
|
||||
await wrapper.get('input[type="password"]').setValue('new-secret')
|
||||
wrapper.getComponent(ProviderPresetSelector).vm.$emit('update:modelValue', '')
|
||||
await flushPromises()
|
||||
expect((wrapper.get('input[type="password"]').element as HTMLInputElement).value).toBe('')
|
||||
})
|
||||
|
||||
it('allocates different credential IDs for two new providers using the same preset', async () => {
|
||||
for (let i = 0; i < 2; i++) {
|
||||
const wrapper = await render()
|
||||
await wrapper.get('[data-preset="deepseek"]').trigger('click')
|
||||
await wrapper.get('input[type="password"]').setValue(`test-key-${i}`)
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
}
|
||||
const ids = vi.mocked(service.putCredential).mock.calls.map(call => call[0])
|
||||
expect(ids).toHaveLength(2)
|
||||
expect(new Set(ids).size).toBe(2)
|
||||
ids.forEach(id => expect(id).toMatch(/^provider-key-[0-9a-f-]{36}$/))
|
||||
vi.mocked(service.createProvider).mock.calls.forEach(([data], index) => {
|
||||
expect(data.credential_id).toBe(ids[index])
|
||||
expect(JSON.stringify(data)).not.toContain('test-key')
|
||||
expect(JSON.stringify(data)).not.toContain('shared-deepseek')
|
||||
})
|
||||
})
|
||||
|
||||
it('accepts a custom API key and clears it after a failed credential save', async () => {
|
||||
const wrapper = await render()
|
||||
await wrapper.get('[data-field="name"]').setValue('Custom')
|
||||
await wrapper.get('[data-field="base-url"]').setValue('https://custom.example.test/v1')
|
||||
await wrapper.get('input[type="password"]').setValue('custom-test-key')
|
||||
vi.mocked(service.putCredential).mockRejectedValueOnce(new Error('credential store unavailable'))
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect((wrapper.get('input[type="password"]').element as HTMLInputElement).value).toBe('')
|
||||
expect(wrapper.text()).toContain('credential store unavailable')
|
||||
expect(service.createProvider).not.toHaveBeenCalled()
|
||||
await wrapper.get('input[type="password"]').setValue('custom-test-key')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(service.createProvider).toHaveBeenCalledWith(expect.objectContaining({ name: 'Custom', credential_id: expect.stringMatching(/^provider-key-/) }))
|
||||
})
|
||||
|
||||
it('preserves its own existing key when untouched and rotates shared legacy references when replacing a key', async () => {
|
||||
const untouched = await render(existing)
|
||||
await untouched.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(service.updateProvider).toHaveBeenLastCalledWith('p1', expect.objectContaining({ credential_id: 'old-shared-key', default_model: 'old-model' }))
|
||||
const rotated = await render(existing)
|
||||
await rotated.get('input[type="password"]').setValue('replacement-test-key')
|
||||
await rotated.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(service.putCredential).toHaveBeenCalledWith(expect.stringMatching(/^provider-key-/), 'replacement-test-key')
|
||||
expect(service.updateProvider).toHaveBeenLastCalledWith('p1', expect.objectContaining({ credential_id: vi.mocked(service.putCredential).mock.calls[0][0] }))
|
||||
})
|
||||
|
||||
it('persists edited protocols and unlinks the previous credential and model', async () => {
|
||||
const wrapper = await render(existing)
|
||||
await wrapper.get('[data-field="protocol"]').setValue('openai_responses')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(service.updateProvider).toHaveBeenCalledWith('p1', expect.objectContaining({ provider_type: 'openai_responses', default_model: '', credential_id: null }))
|
||||
})
|
||||
|
||||
it('does not let a late credential status reuse a key after switching presets', async () => {
|
||||
let resolveStatus!: (configured: boolean) => void
|
||||
vi.mocked(service.getCredentialStatus).mockReturnValue(new Promise(resolve => { resolveStatus = resolve }))
|
||||
const wrapper = await render(existing)
|
||||
await wrapper.get('[data-preset="qwen"]').trigger('click')
|
||||
resolveStatus(true)
|
||||
await flushPromises()
|
||||
await wrapper.get('form').trigger('submit')
|
||||
expect(service.updateProvider).not.toHaveBeenCalled()
|
||||
expect(wrapper.text()).toContain('请输入 API Key')
|
||||
})
|
||||
|
||||
it('clears secrets on close and unmount, and stops a pending credential save from creating a provider', async () => {
|
||||
let finish!: () => void
|
||||
vi.mocked(service.putCredential).mockReturnValue(new Promise(resolve => { finish = resolve }))
|
||||
const wrapper = await render()
|
||||
await wrapper.get('[data-preset="deepseek"]').trigger('click')
|
||||
await wrapper.get('input[type="password"]').setValue('pending-test-key')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await wrapper.get('[aria-label="关闭提供商表单"]').trigger('click')
|
||||
expect((wrapper.get('input[type="password"]').element as HTMLInputElement).value).toBe('')
|
||||
wrapper.unmount()
|
||||
finish()
|
||||
await flushPromises()
|
||||
expect(service.createProvider).not.toHaveBeenCalled()
|
||||
const reopened = await render()
|
||||
expect((reopened.get('input[type="password"]').element as HTMLInputElement).value).toBe('')
|
||||
})
|
||||
|
||||
it('clears a secret after provider save failure and keeps only the successfully saved reference for retry', async () => {
|
||||
const wrapper = await render()
|
||||
await wrapper.get('[data-preset="deepseek"]').trigger('click')
|
||||
await wrapper.get('input[type="password"]').setValue('retry-test-key')
|
||||
vi.mocked(service.createProvider).mockRejectedValueOnce(new Error('provider save failed'))
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(wrapper.text()).toContain('provider save failed')
|
||||
expect((wrapper.get('input[type="password"]').element as HTMLInputElement).value).toBe('')
|
||||
await wrapper.get('form').trigger('submit')
|
||||
await flushPromises()
|
||||
expect(service.putCredential).toHaveBeenCalledTimes(1)
|
||||
expect(service.createProvider).toHaveBeenCalledTimes(2)
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,177 @@
|
||||
<script setup lang="ts">
|
||||
import { computed, nextTick, onBeforeUnmount, onMounted, reactive, ref } from 'vue'
|
||||
import type { ModelInfo, ProviderConfig, ProviderPreset, ProviderType } from '@/contracts'
|
||||
import * as service from '@/services/providerService'
|
||||
import ProviderPresetSelector from './ProviderPresetSelector.vue'
|
||||
|
||||
const props = defineProps<{ provider?: ProviderConfig; models?: ModelInfo[] }>()
|
||||
const emit = defineEmits<{ close: []; saved: [provider: ProviderConfig] }>()
|
||||
const newCredentialId = () => `provider-key-${crypto.randomUUID()}`
|
||||
const form = reactive({
|
||||
preset_id: '', provider_type: props.provider?.provider_type ?? 'openai_compatible' as ProviderType,
|
||||
name: props.provider?.name ?? '', base_url: props.provider?.base_url ?? '',
|
||||
default_model: props.provider?.default_model ?? '', enabled: props.provider?.enabled ?? true,
|
||||
})
|
||||
const credentialId = ref(props.provider?.credential_id || newCredentialId())
|
||||
const apiKey = ref('')
|
||||
const configured = ref(false)
|
||||
const credentialLoading = ref(false)
|
||||
const credentialError = ref('')
|
||||
const presets = ref<ProviderPreset[]>([])
|
||||
const presetsLoading = ref(false)
|
||||
const presetsError = ref('')
|
||||
const saving = ref(false)
|
||||
const error = ref('')
|
||||
const contextChanged = ref(false)
|
||||
const dialog = ref<HTMLElement>()
|
||||
const previousFocus = document.activeElement as HTMLElement | null
|
||||
let active = true
|
||||
let credentialGeneration = 0
|
||||
const selectedPreset = computed(() => presets.value.find(preset => preset.preset_id === form.preset_id))
|
||||
const modelOptions = computed(() => contextChanged.value ? [] : props.models ?? [])
|
||||
|
||||
async function loadPresets() {
|
||||
presetsLoading.value = true
|
||||
presetsError.value = ''
|
||||
try {
|
||||
presets.value = await service.listProviderPresets()
|
||||
if (!contextChanged.value) form.preset_id = presets.value.find(preset => preset.provider_type === props.provider?.provider_type && preset.base_url === props.provider?.base_url)?.preset_id ?? ''
|
||||
} catch { presetsError.value = '预设加载失败,请重试,或填写自定义服务。' }
|
||||
finally { presetsLoading.value = false }
|
||||
}
|
||||
|
||||
onMounted(async () => {
|
||||
void loadPresets()
|
||||
if (props.provider?.credential_id) {
|
||||
const generation = credentialGeneration
|
||||
credentialLoading.value = true
|
||||
try {
|
||||
const result = await service.getCredentialStatus(credentialId.value)
|
||||
if (active && generation === credentialGeneration) configured.value = result
|
||||
} catch {
|
||||
if (active && generation === credentialGeneration) credentialError.value = '无法检查已保存的凭据。可输入新密钥,或关闭后重试。'
|
||||
} finally {
|
||||
if (generation === credentialGeneration) credentialLoading.value = false
|
||||
}
|
||||
}
|
||||
await nextTick()
|
||||
if (active) dialog.value?.querySelector<HTMLInputElement>('input')?.focus()
|
||||
})
|
||||
|
||||
function detachCredential() {
|
||||
credentialGeneration++
|
||||
apiKey.value = ''
|
||||
credentialId.value = newCredentialId()
|
||||
configured.value = false
|
||||
credentialLoading.value = false
|
||||
credentialError.value = ''
|
||||
form.default_model = ''
|
||||
contextChanged.value = true
|
||||
error.value = ''
|
||||
}
|
||||
|
||||
function applyPreset(id: string) {
|
||||
if (form.preset_id === id) return
|
||||
detachCredential()
|
||||
form.preset_id = id
|
||||
const preset = presets.value.find(item => item.preset_id === id)
|
||||
Object.assign(form, { provider_type: preset?.provider_type ?? 'openai_compatible', name: preset?.name ?? '', base_url: preset?.base_url ?? '' })
|
||||
}
|
||||
|
||||
function changeConnection() {
|
||||
form.preset_id = ''
|
||||
detachCredential()
|
||||
}
|
||||
|
||||
function close() {
|
||||
active = false
|
||||
apiKey.value = ''
|
||||
emit('close')
|
||||
}
|
||||
|
||||
onBeforeUnmount(() => {
|
||||
active = false
|
||||
apiKey.value = ''
|
||||
previousFocus?.focus()
|
||||
})
|
||||
|
||||
function handleKeydown(event: KeyboardEvent) {
|
||||
if (event.key === 'Escape') { event.preventDefault(); close() }
|
||||
if (event.key !== 'Tab') return
|
||||
const elements = Array.from(dialog.value?.querySelectorAll<HTMLElement>('button, input, select, [tabindex="0"]') ?? []).filter(element => !element.matches(':disabled'))
|
||||
const first = elements[0], last = elements[elements.length - 1]
|
||||
if (event.shiftKey && document.activeElement === first) { event.preventDefault(); last?.focus() }
|
||||
else if (!event.shiftKey && document.activeElement === last) { event.preventDefault(); first?.focus() }
|
||||
}
|
||||
|
||||
async function save() {
|
||||
if (saving.value || credentialLoading.value || !active) return
|
||||
error.value = ''
|
||||
saving.value = true
|
||||
try {
|
||||
if (!form.name.trim() || (form.provider_type !== 'mock' && !form.base_url.trim())) throw new Error('请填写名称和 Base URL。')
|
||||
if (selectedPreset.value?.requires_credential && !apiKey.value.trim() && !configured.value) throw new Error('请输入 API Key。密钥将由后端加密保存。')
|
||||
// Snapshot before awaiting: closing/unmounting must never create a provider with a changed draft.
|
||||
const data = { provider_type: form.provider_type, name: form.name.trim(), base_url: form.base_url.trim() || undefined, default_model: form.default_model.trim(), enabled: form.enabled, capabilities: {}, has_credential: false }
|
||||
if (apiKey.value.trim()) {
|
||||
// Rotate even an existing reference: older installations may share preset credential IDs.
|
||||
const nextId = newCredentialId()
|
||||
const request = service.putCredential(nextId, apiKey.value.trim())
|
||||
apiKey.value = ''
|
||||
await request
|
||||
if (!active) return
|
||||
credentialId.value = nextId
|
||||
configured.value = true
|
||||
}
|
||||
const reference = configured.value ? credentialId.value : undefined
|
||||
// A failed status check must not silently unlink the provider's existing credential.
|
||||
if (credentialError.value && !reference) throw new Error(credentialError.value)
|
||||
const saved = props.provider
|
||||
? await service.updateProvider(props.provider.provider_id, { ...data, credential_id: reference ?? null })
|
||||
: await service.createProvider({ ...data, credential_id: reference })
|
||||
if (active) { emit('saved', saved); close() }
|
||||
} catch (reason) {
|
||||
if (active) error.value = reason instanceof Error ? reason.message : 'Provider 保存失败,请重试。'
|
||||
} finally { apiKey.value = ''; saving.value = false }
|
||||
}
|
||||
</script>
|
||||
|
||||
<template>
|
||||
<div class="modal-backdrop provider-backdrop" @click.self="close" @keydown="handleKeydown">
|
||||
<div ref="dialog" class="modal provider-modal" role="dialog" aria-modal="true" aria-labelledby="provider-form-title" :aria-busy="saving">
|
||||
<div class="form-heading"><h2 id="provider-form-title">{{ provider ? '编辑 Provider' : '新增 Provider' }}</h2><button type="button" class="button-secondary" aria-label="关闭提供商表单" @click="close">关闭</button></div>
|
||||
<p v-if="presetsLoading" class="subtle" role="status">正在加载提供商预设…</p>
|
||||
<div v-if="presetsError" class="error-banner" role="alert">{{ presetsError }} <button type="button" class="button-secondary" :disabled="presetsLoading || saving" @click="loadPresets">重试</button></div>
|
||||
<form @submit.prevent="save">
|
||||
<fieldset :disabled="saving">
|
||||
<ProviderPresetSelector :presets="presets" :model-value="form.preset_id" @update:model-value="applyPreset" />
|
||||
<p v-if="selectedPreset?.description" class="subtle">{{ selectedPreset.description }}</p>
|
||||
<div class="form-grid">
|
||||
<label class="field"><span>接入协议</span><select v-model="form.provider_type" class="select" data-field="protocol" @change="changeConnection"><option value="openai_compatible">OpenAI Compatible</option><option value="openai_chat">OpenAI Chat</option><option value="openai_responses">OpenAI Responses</option><option value="anthropic_messages">Anthropic Messages</option><option value="ollama">Ollama</option><option v-if="provider?.provider_type === 'mock'" value="mock">Mock</option></select></label>
|
||||
<label class="field"><span>名称</span><input v-model="form.name" class="input" data-field="name" required /></label>
|
||||
<label class="field wide"><span>Base URL</span><input v-model="form.base_url" class="input" data-field="base-url" placeholder="https://api.example.com/v1" :required="form.provider_type !== 'mock'" @change="changeConnection" /></label>
|
||||
<label class="field wide"><span>API Key</span><input v-model="apiKey" class="input" type="password" autocomplete="new-password" spellcheck="false" :placeholder="configured ? '已配置,留空表示不修改' : '请输入 API Key(无鉴权服务可留空)'" /><small class="subtle">密钥由本地 AI Core 加密保存;提供商配置仅保存独立的凭据引用。</small></label>
|
||||
<p v-if="credentialLoading" class="subtle wide" role="status">正在检查凭据状态…</p>
|
||||
<p v-if="credentialError" class="error-text wide" role="alert">{{ credentialError }}</p>
|
||||
<label class="field wide"><span>默认聊天模型</span><input v-model="form.default_model" class="input" data-field="model" list="provider-model-options" placeholder="输入模型 ID,或保存后获取模型列表" /><datalist id="provider-model-options"><option v-for="model in modelOptions" :key="model.model_id" :value="model.model_id">{{ model.name }}</option></datalist></label>
|
||||
</div>
|
||||
<label class="inline-actions"><input v-model="form.enabled" type="checkbox" /> 启用</label>
|
||||
</fieldset>
|
||||
<div v-if="error" class="error-banner" role="alert">{{ error }}</div>
|
||||
<div class="inline-actions form-footer"><button class="button-primary" type="submit" :disabled="saving || credentialLoading">{{ saving ? '保存中…' : '保存提供商' }}</button><button type="button" class="button-secondary" @click="close">取消</button></div>
|
||||
</form>
|
||||
</div>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
<style scoped>
|
||||
.provider-modal { width: min(820px, 100%); max-height: 90dvh; }
|
||||
.form-heading { display: flex; align-items: center; justify-content: space-between; gap: var(--space-md); margin-bottom: var(--space-md); }
|
||||
.form-heading h2 { margin: 0; }
|
||||
fieldset { display: grid; gap: var(--space-md); border: 0; padding: 0; margin: 0; min-width: 0; }
|
||||
.form-grid { display: grid; grid-template-columns: 1fr 1fr; gap: var(--space-md); }
|
||||
.wide { grid-column: 1 / -1; }
|
||||
.error-text { color: var(--color-error); }
|
||||
.form-footer { padding-top: var(--space-sm); }
|
||||
@media (max-width: 600px) { .provider-backdrop { padding: 12px; }.provider-modal { padding: var(--space-lg); max-height: 94dvh; }.form-grid { grid-template-columns: 1fr; } }
|
||||
</style>
|
||||