fix(backend): 修复全面审阅发现的核心问题
修复索引首次失败回滚、Vault 扫描边界、Markdown 代码围栏和 UTF-16 Citation 偏移。 收紧 Plugin 权限与 JSON Schema 校验,补齐 Note Move、Task、Attachment 和 Transcript Tool。 接入 OpenAI SSE 与 Ollama JSONL 真流式输出,修正 Provider PATCH 语义并限制运行时内存保留。 新增对应回归测试,后端测试增至 62 项。
This commit is contained in:
@@ -1,9 +1,10 @@
|
||||
from app.agent.permissions import PermissionManager, PermissionMode, PermissionPolicy
|
||||
from app.agent.runtime import AgentRuntime, AgentRunNotFoundError
|
||||
from app.agent.runtime import AgentCapacityError, AgentRuntime, AgentRunNotFoundError
|
||||
from app.agent.tools import ToolRegistry
|
||||
|
||||
__all__ = [
|
||||
"AgentRunNotFoundError",
|
||||
"AgentCapacityError",
|
||||
"AgentRuntime",
|
||||
"PermissionManager",
|
||||
"PermissionMode",
|
||||
|
||||
@@ -1,9 +1,12 @@
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from app.agent.tools import ToolExecutionContext, ToolRegistry
|
||||
from app.contracts import SearchMode, SearchRequest, ToolDefinition
|
||||
from app.contracts import SearchMode, SearchRequest, TaskStatus, ToolDefinition
|
||||
from app.retrieval.engine import engine
|
||||
from app.services import note_service
|
||||
from app.services import attachment_service, task_service, transcription_service
|
||||
|
||||
|
||||
class ToolArguments(BaseModel):
|
||||
@@ -54,6 +57,42 @@ class NoteListArguments(ToolArguments):
|
||||
tag: str | None = None
|
||||
|
||||
|
||||
class NoteMoveArguments(ToolArguments):
|
||||
note_id: str = Field(min_length=1)
|
||||
folder: str
|
||||
|
||||
|
||||
class TaskCreateArguments(ToolArguments):
|
||||
title: str = Field(min_length=1)
|
||||
description: str = ""
|
||||
note_id: str | None = None
|
||||
due_at: datetime | None = None
|
||||
|
||||
|
||||
class TaskUpdateArguments(ToolArguments):
|
||||
task_id: str = Field(min_length=1)
|
||||
title: str | None = None
|
||||
description: str | None = None
|
||||
status: TaskStatus | None = None
|
||||
note_id: str | None = None
|
||||
due_at: datetime | None = None
|
||||
|
||||
|
||||
class TaskListArguments(ToolArguments):
|
||||
limit: int = Field(default=50, ge=1, le=100)
|
||||
offset: int = Field(default=0, ge=0)
|
||||
|
||||
|
||||
class AttachmentReadArguments(ToolArguments):
|
||||
attachment_id: str = Field(min_length=1)
|
||||
max_chars: int = Field(default=100_000, ge=1, le=1_000_000)
|
||||
|
||||
|
||||
class AudioTranscribeArguments(ToolArguments):
|
||||
attachment_id: str = Field(min_length=1)
|
||||
language: str | None = None
|
||||
|
||||
|
||||
async def echo(arguments: EchoArguments, _: ToolExecutionContext) -> dict[str, str]:
|
||||
return {"text": arguments.text}
|
||||
|
||||
@@ -94,6 +133,39 @@ def list_notes(arguments: NoteListArguments, _: ToolExecutionContext) -> dict:
|
||||
}
|
||||
|
||||
|
||||
async def move_note(arguments: NoteMoveArguments, _: ToolExecutionContext) -> dict:
|
||||
note = await note_service.move_note(arguments.note_id, folder=arguments.folder)
|
||||
return note.model_dump(mode="json")
|
||||
|
||||
|
||||
def create_task(arguments: TaskCreateArguments, _: ToolExecutionContext) -> dict:
|
||||
return task_service.create_task(**arguments.model_dump()).model_dump(mode="json")
|
||||
|
||||
|
||||
def update_task(arguments: TaskUpdateArguments, _: ToolExecutionContext) -> dict:
|
||||
values = arguments.model_dump(exclude_unset=True)
|
||||
task_id = values.pop("task_id")
|
||||
return task_service.update_task(task_id, values).model_dump(mode="json")
|
||||
|
||||
|
||||
def list_tasks(arguments: TaskListArguments, _: ToolExecutionContext) -> dict:
|
||||
items, total = task_service.list_tasks(**arguments.model_dump())
|
||||
return {
|
||||
"items": [item.model_dump(mode="json") for item in items],
|
||||
"page": {"total": total, "limit": arguments.limit, "offset": arguments.offset},
|
||||
}
|
||||
|
||||
|
||||
def read_attachment(arguments: AttachmentReadArguments, _: ToolExecutionContext) -> dict:
|
||||
return attachment_service.read_attachment(**arguments.model_dump())
|
||||
|
||||
|
||||
def transcribe_audio(arguments: AudioTranscribeArguments, _: ToolExecutionContext) -> dict:
|
||||
return transcription_service.create_transcription(
|
||||
arguments.attachment_id, arguments.language
|
||||
).model_dump(mode="json")
|
||||
|
||||
|
||||
def _register(
|
||||
registry: ToolRegistry,
|
||||
*,
|
||||
@@ -178,3 +250,51 @@ def register_builtin_tools(registry: ToolRegistry) -> None:
|
||||
executor=list_notes,
|
||||
permission="notes.read",
|
||||
)
|
||||
_register(
|
||||
registry,
|
||||
name="notes.move",
|
||||
description="Move a note to another folder while preserving note_id.",
|
||||
arguments_model=NoteMoveArguments,
|
||||
executor=move_note,
|
||||
permission="notes.write",
|
||||
)
|
||||
_register(
|
||||
registry,
|
||||
name="tasks.create",
|
||||
description="Create a persistent task.",
|
||||
arguments_model=TaskCreateArguments,
|
||||
executor=create_task,
|
||||
permission="tasks.write",
|
||||
)
|
||||
_register(
|
||||
registry,
|
||||
name="tasks.update",
|
||||
description="Update a persistent task.",
|
||||
arguments_model=TaskUpdateArguments,
|
||||
executor=update_task,
|
||||
permission="tasks.write",
|
||||
)
|
||||
_register(
|
||||
registry,
|
||||
name="tasks.list",
|
||||
description="List persistent tasks.",
|
||||
arguments_model=TaskListArguments,
|
||||
executor=list_tasks,
|
||||
permission="tasks.read",
|
||||
)
|
||||
_register(
|
||||
registry,
|
||||
name="attachments.read",
|
||||
description="Read a UTF-8 attachment from host-managed attachment storage.",
|
||||
arguments_model=AttachmentReadArguments,
|
||||
executor=read_attachment,
|
||||
permission="attachments.read",
|
||||
)
|
||||
_register(
|
||||
registry,
|
||||
name="audio.transcribe",
|
||||
description="Read a host-generated transcript for an audio attachment.",
|
||||
arguments_model=AudioTranscribeArguments,
|
||||
executor=transcribe_audio,
|
||||
permission="attachments.read",
|
||||
)
|
||||
|
||||
@@ -10,13 +10,39 @@ class PermissionMode(str, Enum):
|
||||
deny = "deny"
|
||||
|
||||
|
||||
KNOWN_PERMISSIONS = frozenset(
|
||||
{
|
||||
"notes.read",
|
||||
"notes.search",
|
||||
"notes.write",
|
||||
"notes.delete",
|
||||
"tasks.read",
|
||||
"tasks.write",
|
||||
"attachments.read",
|
||||
"network.request",
|
||||
"secrets.use",
|
||||
"ui.command",
|
||||
"ui.settings",
|
||||
"ui.sidebar",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class PermissionPolicy:
|
||||
def __init__(self) -> None:
|
||||
self._rules: dict[str, PermissionMode] = {
|
||||
"notes.read": PermissionMode.allow,
|
||||
"notes.search": PermissionMode.allow,
|
||||
"notes.delete": PermissionMode.confirm,
|
||||
"notes.write": PermissionMode.confirm,
|
||||
"tasks.read": PermissionMode.allow,
|
||||
"tasks.write": PermissionMode.confirm,
|
||||
"attachments.read": PermissionMode.allow,
|
||||
"network.request": PermissionMode.confirm,
|
||||
"secrets.use": PermissionMode.confirm,
|
||||
"ui.command": PermissionMode.allow,
|
||||
"ui.settings": PermissionMode.allow,
|
||||
"ui.sidebar": PermissionMode.allow,
|
||||
}
|
||||
|
||||
def set_rule(self, permission: str, mode: PermissionMode) -> None:
|
||||
@@ -25,7 +51,7 @@ class PermissionPolicy:
|
||||
def mode_for(self, permission: str | None) -> PermissionMode:
|
||||
if permission is None:
|
||||
return PermissionMode.allow
|
||||
return self._rules.get(permission, PermissionMode.allow)
|
||||
return self._rules.get(permission, PermissionMode.deny)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
|
||||
@@ -34,11 +34,18 @@ class AgentRunNotFoundError(LookupError):
|
||||
pass
|
||||
|
||||
|
||||
class AgentCapacityError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
TERMINAL_STATUSES = {
|
||||
AgentRunStatus.completed,
|
||||
AgentRunStatus.failed,
|
||||
AgentRunStatus.cancelled,
|
||||
}
|
||||
MAX_RUN_RECORDS = 200
|
||||
MAX_EVENTS_PER_RUN = 2_000
|
||||
MAX_TOOL_CALLS_PER_TURN = 50
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
@@ -67,6 +74,7 @@ class AgentRuntime:
|
||||
self._records: dict[str, RunRecord] = {}
|
||||
|
||||
async def create_run(self, request: AgentRunCreateRequest) -> AgentRun:
|
||||
self._prune_records()
|
||||
provider = self.providers.get(request.provider_id)
|
||||
skill_config = None
|
||||
if request.skill_id:
|
||||
@@ -217,6 +225,13 @@ class AgentRuntime:
|
||||
return
|
||||
|
||||
if turn.tool_calls:
|
||||
if len(turn.tool_calls) > MAX_TOOL_CALLS_PER_TURN:
|
||||
self._fail(
|
||||
record,
|
||||
"TOO_MANY_TOOL_CALLS",
|
||||
f"Provider requested more than {MAX_TOOL_CALLS_PER_TURN} tools in one turn.",
|
||||
)
|
||||
return
|
||||
calls = [
|
||||
ToolCall(
|
||||
tool_call_id=item.tool_call_id,
|
||||
@@ -393,6 +408,8 @@ class AgentRuntime:
|
||||
timestamp=datetime.now(timezone.utc),
|
||||
)
|
||||
record.events.append(event)
|
||||
if len(record.events) > MAX_EVENTS_PER_RUN:
|
||||
del record.events[: len(record.events) - MAX_EVENTS_PER_RUN]
|
||||
for queue in record.subscribers:
|
||||
queue.put_nowait(event)
|
||||
|
||||
@@ -429,3 +446,20 @@ class AgentRuntime:
|
||||
return self._records[run_id]
|
||||
except KeyError as exc:
|
||||
raise AgentRunNotFoundError(run_id) from exc
|
||||
|
||||
def _prune_records(self) -> None:
|
||||
overflow = len(self._records) - MAX_RUN_RECORDS + 1
|
||||
if overflow <= 0:
|
||||
return
|
||||
terminal = sorted(
|
||||
(
|
||||
record
|
||||
for record in self._records.values()
|
||||
if record.run.status in TERMINAL_STATUSES
|
||||
),
|
||||
key=lambda record: record.run.updated_at,
|
||||
)
|
||||
for record in terminal[:overflow]:
|
||||
self._records.pop(record.run.run_id, None)
|
||||
if len(self._records) >= MAX_RUN_RECORDS:
|
||||
raise AgentCapacityError("Too many active Agent runs.")
|
||||
|
||||
@@ -4,6 +4,8 @@ from time import perf_counter
|
||||
from typing import Any, Awaitable, Callable
|
||||
|
||||
from pydantic import BaseModel, ValidationError
|
||||
from jsonschema import Draft202012Validator
|
||||
from jsonschema.exceptions import ValidationError as JsonSchemaValidationError
|
||||
|
||||
from app.contracts import ToolCall, ToolDefinition, ToolResult
|
||||
|
||||
@@ -78,8 +80,9 @@ class ToolRegistry:
|
||||
)
|
||||
|
||||
try:
|
||||
Draft202012Validator(registered.definition.parameters).validate(call.arguments)
|
||||
arguments = registered.arguments_model.model_validate(call.arguments)
|
||||
except ValidationError as exc:
|
||||
except (ValidationError, JsonSchemaValidationError) as exc:
|
||||
return ToolResult(
|
||||
tool_call_id=call.tool_call_id,
|
||||
name=call.name,
|
||||
|
||||
@@ -23,6 +23,7 @@ class Settings:
|
||||
data_dir: Path
|
||||
db_path: Path
|
||||
vault_path: Path
|
||||
attachments_path: Path
|
||||
|
||||
|
||||
@lru_cache
|
||||
@@ -37,4 +38,7 @@ def get_settings() -> Settings:
|
||||
data_dir=data_dir,
|
||||
db_path=Path(os.getenv("APP_DB_PATH", str(data_dir / "app.db"))),
|
||||
vault_path=Path(os.getenv("APP_VAULT_PATH", str(data_dir / "vault"))),
|
||||
attachments_path=Path(
|
||||
os.getenv("APP_ATTACHMENTS_PATH", str(data_dir / "attachments"))
|
||||
),
|
||||
)
|
||||
|
||||
@@ -384,6 +384,7 @@ class Plugin(Contract):
|
||||
manifest: PluginManifest
|
||||
status: PluginStatus
|
||||
enabled: bool = False
|
||||
granted_permissions: list[str] = Field(default_factory=list)
|
||||
error_message: str | None = None
|
||||
|
||||
|
||||
@@ -391,6 +392,10 @@ class PluginListResponse(Contract):
|
||||
items: list[Plugin] = Field(default_factory=list)
|
||||
|
||||
|
||||
class PluginPermissionGrantRequest(Contract):
|
||||
permissions: list[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
# Providers
|
||||
class ProviderType(str, Enum):
|
||||
mock = "mock"
|
||||
@@ -506,6 +511,9 @@ class TranscriptionJob(Contract):
|
||||
job_id: str
|
||||
attachment_id: str
|
||||
status: Literal["queued", "processing", "completed", "failed"]
|
||||
text: str | None = None
|
||||
error_code: str | None = None
|
||||
error_message: str | None = None
|
||||
created_at: datetime
|
||||
|
||||
|
||||
|
||||
@@ -55,6 +55,20 @@ MIGRATIONS: list[str] = [
|
||||
value TEXT NOT NULL
|
||||
);
|
||||
""",
|
||||
# v2: 第一阶段 Task Core;正文仍归 Note/Vault,任务状态持久化到 SQLite。
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS tasks (
|
||||
task_id TEXT PRIMARY KEY,
|
||||
title TEXT NOT NULL,
|
||||
description TEXT NOT NULL DEFAULT '',
|
||||
status TEXT NOT NULL DEFAULT 'todo',
|
||||
note_id TEXT REFERENCES notes(note_id) ON DELETE SET NULL,
|
||||
due_at TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_tasks_status_due ON tasks(status, due_at);
|
||||
""",
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -17,6 +17,7 @@ class ApiError(Exception):
|
||||
message: str,
|
||||
details: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
super().__init__(message)
|
||||
self.status_code = status_code
|
||||
self.code = code
|
||||
self.message = message
|
||||
|
||||
@@ -6,9 +6,12 @@ from pathlib import Path
|
||||
from typing import Any, Literal
|
||||
|
||||
import yaml
|
||||
from jsonschema import Draft202012Validator
|
||||
from jsonschema.exceptions import SchemaError
|
||||
from pydantic import BaseModel, ConfigDict, Field, ValidationError, create_model
|
||||
|
||||
from app.agent.tools import ToolExecutionContext, ToolRegistry
|
||||
from app.agent.permissions import KNOWN_PERMISSIONS
|
||||
from app.contracts import (
|
||||
ModelCapability,
|
||||
Plugin,
|
||||
@@ -73,6 +76,7 @@ class SkillRuntime:
|
||||
except ValidationError as exc:
|
||||
raise _manifest_error("skill", exc) from exc
|
||||
_validate_id("skill", manifest.skill_id)
|
||||
_validate_permissions("skill", manifest.permissions)
|
||||
if manifest.skill_id in self._records:
|
||||
raise ExtensionError(
|
||||
"SKILL_ALREADY_INSTALLED",
|
||||
@@ -257,6 +261,7 @@ class PluginRuntime:
|
||||
except ValidationError as exc:
|
||||
raise _manifest_error("plugin", exc) from exc
|
||||
_validate_id("plugin", manifest.plugin_id)
|
||||
_validate_permissions("plugin", manifest.permissions)
|
||||
if manifest.plugin_id in self._records:
|
||||
raise ExtensionError(
|
||||
"PLUGIN_ALREADY_INSTALLED",
|
||||
@@ -275,6 +280,7 @@ class PluginRuntime:
|
||||
)
|
||||
for spec in specs:
|
||||
_validate_id("tool", spec.name)
|
||||
_validate_tool_schema(spec)
|
||||
if spec.permission and spec.permission not in manifest.permissions:
|
||||
raise ExtensionError(
|
||||
"PLUGIN_PERMISSION_UNDECLARED",
|
||||
@@ -283,7 +289,14 @@ class PluginRuntime:
|
||||
)
|
||||
|
||||
record = _PluginRecord(
|
||||
plugin=Plugin(manifest=manifest, status=PluginStatus.installed),
|
||||
plugin=Plugin(
|
||||
manifest=manifest,
|
||||
status=(
|
||||
PluginStatus.permission_required
|
||||
if manifest.permissions
|
||||
else PluginStatus.installed
|
||||
),
|
||||
),
|
||||
tools=specs,
|
||||
package_path=root,
|
||||
registered_tools=[],
|
||||
@@ -309,6 +322,17 @@ class PluginRuntime:
|
||||
status_code=501,
|
||||
details={"plugin_id": plugin_id, "backend": "mcp"},
|
||||
)
|
||||
missing_grants = sorted(
|
||||
set(record.plugin.manifest.permissions) - set(record.plugin.granted_permissions)
|
||||
)
|
||||
if missing_grants:
|
||||
record.plugin.status = PluginStatus.permission_required
|
||||
raise ExtensionError(
|
||||
"PLUGIN_PERMISSION_REQUIRED",
|
||||
"Plugin permissions must be granted before it can be enabled.",
|
||||
status_code=409,
|
||||
details={"plugin_id": plugin_id, "permissions": missing_grants},
|
||||
)
|
||||
conflicts = [spec.name for spec in record.tools if self.registry.contains(spec.name)]
|
||||
if conflicts:
|
||||
raise ExtensionError(
|
||||
@@ -353,6 +377,27 @@ class PluginRuntime:
|
||||
record.plugin.error_message = None
|
||||
return record.plugin.model_copy(deep=True)
|
||||
|
||||
def set_permissions(self, plugin_id: str, permissions: list[str]) -> Plugin:
|
||||
record = self._record(plugin_id)
|
||||
requested = set(permissions)
|
||||
declared = set(record.plugin.manifest.permissions)
|
||||
undeclared = sorted(requested - declared)
|
||||
if undeclared:
|
||||
raise ExtensionError(
|
||||
"PLUGIN_PERMISSION_UNDECLARED",
|
||||
"Cannot grant permissions that are not declared by the Plugin.",
|
||||
details={"plugin_id": plugin_id, "permissions": undeclared},
|
||||
)
|
||||
record.plugin.granted_permissions = sorted(requested)
|
||||
missing = declared - requested
|
||||
if missing and record.plugin.enabled:
|
||||
self.disable(plugin_id)
|
||||
if missing:
|
||||
record.plugin.status = PluginStatus.permission_required
|
||||
elif not record.plugin.enabled:
|
||||
record.plugin.status = PluginStatus.installed
|
||||
return record.plugin.model_copy(deep=True)
|
||||
|
||||
def disable(self, plugin_id: str) -> Plugin:
|
||||
record = self._record(plugin_id)
|
||||
for name in record.registered_tools:
|
||||
@@ -429,6 +474,16 @@ def _validate_id(kind: str, value: str) -> None:
|
||||
)
|
||||
|
||||
|
||||
def _validate_permissions(kind: str, permissions: list[str]) -> None:
|
||||
unknown = sorted(set(permissions) - KNOWN_PERMISSIONS)
|
||||
if unknown:
|
||||
raise ExtensionError(
|
||||
"EXTENSION_PERMISSION_INVALID",
|
||||
f"Invalid {kind} permissions: {', '.join(unknown)}",
|
||||
details={"kind": kind, "permissions": unknown},
|
||||
)
|
||||
|
||||
|
||||
def _manifest_error(kind: str, exc: ValidationError) -> ExtensionError:
|
||||
return ExtensionError(
|
||||
"EXTENSION_MANIFEST_INVALID",
|
||||
@@ -457,3 +512,23 @@ def _arguments_model(spec: DeclarativeToolSpec) -> type[BaseModel]:
|
||||
fields[name] = (annotation, ... if name in required else None)
|
||||
model_name = "PluginArgs_" + re.sub(r"\W+", "_", spec.name)
|
||||
return create_model(model_name, __config__=ConfigDict(extra="forbid"), **fields)
|
||||
|
||||
|
||||
def _validate_tool_schema(spec: DeclarativeToolSpec) -> None:
|
||||
schema = spec.parameters or {"type": "object", "properties": {}}
|
||||
try:
|
||||
Draft202012Validator.check_schema(schema)
|
||||
except SchemaError as exc:
|
||||
raise ExtensionError(
|
||||
"PLUGIN_TOOL_SCHEMA_INVALID",
|
||||
f"Invalid JSON Schema for tool {spec.name}: {exc.message}",
|
||||
details={"tool": spec.name},
|
||||
) from exc
|
||||
if schema.get("type", "object") != "object" or not isinstance(
|
||||
schema.get("properties", {}), dict
|
||||
):
|
||||
raise ExtensionError(
|
||||
"PLUGIN_TOOL_SCHEMA_INVALID",
|
||||
"Tool parameters must be an object schema with object properties.",
|
||||
details={"tool": spec.name},
|
||||
)
|
||||
|
||||
@@ -18,6 +18,7 @@ from app.textutils import count_tokens
|
||||
|
||||
_HEADING_RE = re.compile(r"^(#{1,6})[ \t]+(.*?)\s*$")
|
||||
_FRONTMATTER_KEY_RE = re.compile(r"^([A-Za-z0-9_-]+)\s*:\s*(.*)$")
|
||||
_FENCE_RE = re.compile(r"^[ \t]{0,3}(`{3,}|~{3,})(?:[^`]*)$")
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -48,9 +49,10 @@ def parse_note(
|
||||
tags: list[str] | None = None,
|
||||
created_at: datetime,
|
||||
updated_at: datetime,
|
||||
note_id: str | None = None,
|
||||
) -> ParsedNote:
|
||||
"""解析一篇 Markdown,生成 ParsedNote(元数据 + Block 列表)。"""
|
||||
note_id = note_id_for_path(file_path)
|
||||
note_id = note_id or note_id_for_path(file_path)
|
||||
frontmatter = _extract_frontmatter(markdown)
|
||||
|
||||
fallback_title = Path(file_path).stem
|
||||
@@ -79,13 +81,14 @@ def parse_blocks(markdown: str, note_id: str) -> list[NoteBlock]:
|
||||
heading_stack: list[str] = []
|
||||
body: list[tuple[str, int]] = []
|
||||
id_counters: dict[str, int] = {}
|
||||
fence_marker: str | None = None
|
||||
|
||||
def make_block(path: list[str], chunk: list[tuple[str, int]]) -> None:
|
||||
if not chunk:
|
||||
return
|
||||
content = "\n".join(line for line, _ in chunk)
|
||||
start = chunk[0][1]
|
||||
end = chunk[-1][1] + len(chunk[-1][0])
|
||||
end = chunk[-1][1] + _utf16_len(chunk[-1][0])
|
||||
block_id = _stable_block_id(note_id, path, content, id_counters)
|
||||
blocks.append(
|
||||
NoteBlock(
|
||||
@@ -109,6 +112,20 @@ def parse_blocks(markdown: str, note_id: str) -> list[NoteBlock]:
|
||||
if offset < content_start:
|
||||
continue # 跳过 frontmatter 区域,但保留 offset 准确性
|
||||
|
||||
fence = _FENCE_RE.match(line)
|
||||
if fence_marker is not None:
|
||||
body.append((line, offset))
|
||||
marker = fence.group(1) if fence else ""
|
||||
if marker.startswith(fence_marker[0]) and len(marker) >= len(fence_marker):
|
||||
fence_marker = None
|
||||
flush_body()
|
||||
continue
|
||||
if fence:
|
||||
flush_body()
|
||||
fence_marker = fence.group(1)
|
||||
body.append((line, offset))
|
||||
continue
|
||||
|
||||
heading = _HEADING_RE.match(line)
|
||||
if heading:
|
||||
flush_body()
|
||||
@@ -138,7 +155,7 @@ def _stable_block_id(note_id: str, path: list[str], content: str, counters: dict
|
||||
|
||||
|
||||
def _split_lines(text: str) -> list[tuple[str, int]]:
|
||||
"""按行拆分并记录每行在原文中的起始字符偏移。"""
|
||||
"""按行拆分并记录 UTF-16 code unit 偏移,直接兼容浏览器编辑器。"""
|
||||
result: list[tuple[str, int]] = []
|
||||
start = 0
|
||||
for raw in text.splitlines(keepends=True):
|
||||
@@ -148,19 +165,23 @@ def _split_lines(text: str) -> list[tuple[str, int]]:
|
||||
elif line.endswith("\n") or line.endswith("\r"):
|
||||
line = line[:-1]
|
||||
result.append((line, start))
|
||||
start += len(raw)
|
||||
start += _utf16_len(raw)
|
||||
return result
|
||||
|
||||
|
||||
def _content_start(markdown: str) -> int:
|
||||
"""返回正文起始偏移:有 frontmatter 时跳过 --- 分隔块,否则为 0。"""
|
||||
"""返回正文起始 UTF-16 偏移:有 frontmatter 时跳过 --- 分隔块。"""
|
||||
if markdown.startswith("---"):
|
||||
end = markdown.find("\n---", 3)
|
||||
if end != -1:
|
||||
return end + 4
|
||||
return _utf16_len(markdown[: end + 4])
|
||||
return 0
|
||||
|
||||
|
||||
def _utf16_len(text: str) -> int:
|
||||
return len(text.encode("utf-16-le")) // 2
|
||||
|
||||
|
||||
def _extract_frontmatter(markdown: str) -> dict[str, str]:
|
||||
"""极简 frontmatter 解析,只提取 key: value 行。"""
|
||||
if not markdown.startswith("---"):
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
from uuid import uuid4
|
||||
import json
|
||||
from collections.abc import AsyncIterator
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import httpx
|
||||
|
||||
from app.contracts import ModelCapability, ModelInfo, ModelRequest
|
||||
from app.contracts import ModelCapability, ModelEvent, ModelEventType, ModelInfo, ModelRequest
|
||||
from app.providers.base import ProviderError, ProviderToolCall, ProviderTurn
|
||||
from app.providers.http_base import TurnStreamingMixin, decode_tool_arguments
|
||||
|
||||
@@ -86,6 +89,115 @@ class OllamaProvider(TurnStreamingMixin):
|
||||
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 {}
|
||||
if message.get("thinking"):
|
||||
yield event(ModelEventType.thinking_delta, {"text": 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)
|
||||
|
||||
def _chat_payload(self, request: ModelRequest, *, stream: bool) -> dict[str, object]:
|
||||
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": stream
|
||||
}
|
||||
if request.tools:
|
||||
payload["tools"] = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tool.name,
|
||||
"description": tool.description,
|
||||
"parameters": tool.parameters,
|
||||
},
|
||||
}
|
||||
for tool in request.tools
|
||||
]
|
||||
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 test_connection(self, model: str | None = None) -> tuple[bool, str]:
|
||||
try:
|
||||
models = await self.list_models()
|
||||
|
||||
@@ -1,9 +1,18 @@
|
||||
import json
|
||||
from collections.abc import AsyncIterator
|
||||
from datetime import datetime, timezone
|
||||
from uuid import uuid4
|
||||
|
||||
import httpx
|
||||
|
||||
from app.contracts import MessageRole, ModelCapability, ModelInfo, ModelRequest
|
||||
from app.contracts import (
|
||||
MessageRole,
|
||||
ModelCapability,
|
||||
ModelEvent,
|
||||
ModelEventType,
|
||||
ModelInfo,
|
||||
ModelRequest,
|
||||
)
|
||||
from app.providers.base import ProviderError, ProviderToolCall, ProviderTurn
|
||||
from app.providers.credentials import CredentialResolver
|
||||
from app.providers.http_base import TurnStreamingMixin, decode_tool_arguments
|
||||
@@ -25,29 +34,7 @@ class OpenAICompatibleProvider(TurnStreamingMixin):
|
||||
self.transport = transport
|
||||
|
||||
async def complete(self, request: ModelRequest) -> ProviderTurn:
|
||||
payload: dict[str, object] = {
|
||||
"model": request.model,
|
||||
"messages": self._messages(request),
|
||||
"stream": False,
|
||||
}
|
||||
if request.tools:
|
||||
payload["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
|
||||
if request.max_tokens is not None:
|
||||
payload["max_tokens"] = request.max_tokens
|
||||
if request.response_format is not None:
|
||||
payload["response_format"] = request.response_format
|
||||
payload = self._payload(request, stream=False)
|
||||
|
||||
data = await self._request("POST", "/chat/completions", json=payload)
|
||||
try:
|
||||
@@ -73,6 +60,133 @@ class OpenAICompatibleProvider(TurnStreamingMixin):
|
||||
output_tokens=int(usage.get("completion_tokens") or 0),
|
||||
)
|
||||
|
||||
def _payload(self, request: ModelRequest, *, stream: bool) -> dict[str, object]:
|
||||
payload: dict[str, object] = {
|
||||
"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
|
||||
]
|
||||
if request.temperature is not None:
|
||||
payload["temperature"] = request.temperature
|
||||
if request.max_tokens is not None:
|
||||
payload["max_tokens"] = request.max_tokens
|
||||
if request.response_format is not None:
|
||||
payload["response_format"] = request.response_format
|
||||
|
||||
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 []
|
||||
if not choices:
|
||||
continue
|
||||
choice = choices[0]
|
||||
delta = choice.get("delta") or {}
|
||||
if delta.get("reasoning_content"):
|
||||
yield event(
|
||||
ModelEventType.thinking_delta,
|
||||
{"text": 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
|
||||
|
||||
async def list_models(self) -> list[ModelInfo]:
|
||||
data = await self._request("GET", "/models")
|
||||
return [
|
||||
@@ -127,10 +241,7 @@ class OpenAICompatibleProvider(TurnStreamingMixin):
|
||||
return result
|
||||
|
||||
async def _request(self, method: str, path: str, **kwargs) -> dict:
|
||||
headers = {"Content-Type": "application/json"}
|
||||
api_key = self.credentials.resolve(self.credential_id)
|
||||
if api_key:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
headers = self._headers()
|
||||
try:
|
||||
async with httpx.AsyncClient(
|
||||
timeout=self.timeout_seconds, transport=self.transport
|
||||
@@ -143,14 +254,25 @@ class OpenAICompatibleProvider(TurnStreamingMixin):
|
||||
except httpx.TimeoutException as exc:
|
||||
raise ProviderError("PROVIDER_TIMEOUT", "Provider request timed out.") from exc
|
||||
except httpx.HTTPStatusError as exc:
|
||||
code = {
|
||||
401: "PROVIDER_AUTH_FAILED",
|
||||
404: "MODEL_NOT_FOUND",
|
||||
429: "PROVIDER_RATE_LIMITED",
|
||||
}.get(exc.response.status_code, "PROVIDER_UNAVAILABLE")
|
||||
raise ProviderError(code, f"Provider returned HTTP {exc.response.status_code}.") from 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"}
|
||||
api_key = self.credentials.resolve(self.credential_id)
|
||||
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}.")
|
||||
|
||||
+63
-39
@@ -10,7 +10,6 @@ from app.contracts import (
|
||||
AgentRunCreateRequest,
|
||||
AgentRunListResponse,
|
||||
ChatRequest,
|
||||
ErrorResponse,
|
||||
ExtensionInstallRequest,
|
||||
IndexJob,
|
||||
IndexRebuildRequest,
|
||||
@@ -27,6 +26,7 @@ from app.contracts import (
|
||||
PermissionDecisionRequest,
|
||||
Plugin,
|
||||
PluginListResponse,
|
||||
PluginPermissionGrantRequest,
|
||||
ProviderConfig,
|
||||
ProviderCreateRequest,
|
||||
ProviderListResponse,
|
||||
@@ -46,17 +46,16 @@ from app.contracts import (
|
||||
TranscriptionJob,
|
||||
TranscriptionRequest,
|
||||
)
|
||||
from app.agent import AgentRunNotFoundError
|
||||
from app.agent import AgentCapacityError, AgentRunNotFoundError
|
||||
from app.container import container
|
||||
from app.errors import ApiError, not_implemented
|
||||
from app.errors import ApiError
|
||||
from app.extensions import ExtensionError
|
||||
from app.providers.registry import ProviderNotFoundError
|
||||
from app.providers.factory import UnsupportedProviderError
|
||||
from app.retrieval.engine import engine
|
||||
from app.services import index_service, note_service
|
||||
from app.services import index_service, note_service, task_service, transcription_service
|
||||
|
||||
router = APIRouter(prefix="/api")
|
||||
not_implemented_response = {501: {"model": ErrorResponse, "description": "业务服务尚未实现"}}
|
||||
|
||||
|
||||
def utc_now() -> datetime:
|
||||
@@ -151,11 +150,9 @@ async def delete_note(note_id: str) -> OperationResponse:
|
||||
return OperationResponse(status="completed", resource_id=note_id, message="deleted")
|
||||
|
||||
|
||||
@router.post(
|
||||
"/notes/{note_id}/move", response_model=Note, responses=not_implemented_response, tags=["Notes"]
|
||||
)
|
||||
async def move_note(note_id: str, _: NoteMoveRequest) -> Note:
|
||||
not_implemented(f"notes.move:{note_id}")
|
||||
@router.post("/notes/{note_id}/move", response_model=Note, tags=["Notes"])
|
||||
async def move_note(note_id: str, request: NoteMoveRequest) -> Note:
|
||||
return await note_service.move_note(note_id, folder=request.folder)
|
||||
|
||||
|
||||
# Retrieval and chat
|
||||
@@ -219,6 +216,8 @@ async def create_agent_run(request: AgentRunCreateRequest) -> AgentRun:
|
||||
return await container.agent.create_run(request)
|
||||
except ExtensionError as exc:
|
||||
raise ApiError(exc.status_code, exc.code, exc.message, exc.details) from exc
|
||||
except AgentCapacityError as exc:
|
||||
raise ApiError(429, "AGENT_CAPACITY_EXCEEDED", str(exc)) from exc
|
||||
|
||||
|
||||
@router.get(
|
||||
@@ -386,6 +385,19 @@ async def disable_plugin(plugin_id: str) -> Plugin:
|
||||
return extension_call(lambda: container.plugins.disable(plugin_id))
|
||||
|
||||
|
||||
@router.put(
|
||||
"/plugins/{plugin_id}/permissions",
|
||||
response_model=Plugin,
|
||||
tags=["Plugins"],
|
||||
)
|
||||
async def set_plugin_permissions(
|
||||
plugin_id: str, request: PluginPermissionGrantRequest
|
||||
) -> Plugin:
|
||||
return extension_call(
|
||||
lambda: container.plugins.set_permissions(plugin_id, request.permissions)
|
||||
)
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/plugins/{plugin_id}",
|
||||
response_model=OperationResponse,
|
||||
@@ -407,7 +419,6 @@ async def list_providers() -> ProviderListResponse:
|
||||
@router.get(
|
||||
"/providers/{provider_id}",
|
||||
response_model=ProviderConfig,
|
||||
responses=not_implemented_response,
|
||||
tags=["Providers"],
|
||||
)
|
||||
async def get_provider(provider_id: str) -> ProviderConfig:
|
||||
@@ -417,7 +428,6 @@ async def get_provider(provider_id: str) -> ProviderConfig:
|
||||
@router.post(
|
||||
"/providers",
|
||||
response_model=ProviderConfig,
|
||||
responses=not_implemented_response,
|
||||
tags=["Providers"],
|
||||
)
|
||||
async def create_provider(request: ProviderCreateRequest) -> ProviderConfig:
|
||||
@@ -446,7 +456,6 @@ async def create_provider(request: ProviderCreateRequest) -> ProviderConfig:
|
||||
@router.patch(
|
||||
"/providers/{provider_id}",
|
||||
response_model=ProviderConfig,
|
||||
responses=not_implemented_response,
|
||||
tags=["Providers"],
|
||||
)
|
||||
async def update_provider(
|
||||
@@ -455,7 +464,19 @@ async def update_provider(
|
||||
current = configurable_provider_or_404(provider_id).config
|
||||
if provider_id == "mock":
|
||||
raise ApiError(409, "BUILTIN_PROVIDER_IMMUTABLE", "Mock provider cannot be modified.")
|
||||
config = current.model_copy(update=request.model_dump(exclude_none=True))
|
||||
fields = request.model_fields_set
|
||||
if ("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.",
|
||||
)
|
||||
updates = {name: getattr(request, name) for name in fields}
|
||||
config = ProviderConfig.model_validate(
|
||||
{**current.model_dump(mode="python"), **updates}
|
||||
)
|
||||
adapter = container.provider_factory.build(config)
|
||||
container.providers.replace(config, adapter)
|
||||
return config
|
||||
@@ -464,7 +485,6 @@ async def update_provider(
|
||||
@router.delete(
|
||||
"/providers/{provider_id}",
|
||||
response_model=OperationResponse,
|
||||
responses=not_implemented_response,
|
||||
tags=["Providers"],
|
||||
)
|
||||
async def delete_provider(provider_id: str) -> OperationResponse:
|
||||
@@ -478,7 +498,6 @@ async def delete_provider(provider_id: str) -> OperationResponse:
|
||||
@router.get(
|
||||
"/providers/{provider_id}/models",
|
||||
response_model=ProviderModelsResponse,
|
||||
responses=not_implemented_response,
|
||||
tags=["Providers"],
|
||||
)
|
||||
async def list_provider_models(provider_id: str) -> ProviderModelsResponse:
|
||||
@@ -492,7 +511,6 @@ async def list_provider_models(provider_id: str) -> ProviderModelsResponse:
|
||||
@router.post(
|
||||
"/providers/test",
|
||||
response_model=ProviderTestResponse,
|
||||
responses=not_implemented_response,
|
||||
tags=["Providers"],
|
||||
)
|
||||
async def test_provider(request: ProviderTestRequest) -> ProviderTestResponse:
|
||||
@@ -517,38 +535,39 @@ async def test_provider(request: ProviderTestRequest) -> ProviderTestResponse:
|
||||
async def list_tasks(
|
||||
limit: int = Query(default=50, ge=1, le=100), offset: int = Query(default=0, ge=0)
|
||||
) -> TaskListResponse:
|
||||
return TaskListResponse(page=PageMeta(limit=limit, offset=offset))
|
||||
items, total = task_service.list_tasks(limit=limit, offset=offset)
|
||||
return TaskListResponse(
|
||||
items=items, page=PageMeta(total=total, limit=limit, offset=offset)
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/tasks", response_model=Task, responses=not_implemented_response, tags=["Tasks"]
|
||||
)
|
||||
async def create_task(_: TaskCreateRequest) -> Task:
|
||||
not_implemented("tasks.create")
|
||||
@router.post("/tasks", response_model=Task, tags=["Tasks"])
|
||||
async def create_task(request: TaskCreateRequest) -> Task:
|
||||
return task_service.create_task(**request.model_dump())
|
||||
|
||||
|
||||
@router.get(
|
||||
"/tasks/{task_id}", response_model=Task, responses=not_implemented_response, tags=["Tasks"]
|
||||
)
|
||||
@router.get("/tasks/{task_id}", response_model=Task, tags=["Tasks"])
|
||||
async def get_task(task_id: str) -> Task:
|
||||
not_implemented(f"tasks.read:{task_id}")
|
||||
task = task_service.get_task(task_id)
|
||||
if task is None:
|
||||
raise ApiError(404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id})
|
||||
return task
|
||||
|
||||
|
||||
@router.patch(
|
||||
"/tasks/{task_id}", response_model=Task, responses=not_implemented_response, tags=["Tasks"]
|
||||
)
|
||||
async def update_task(task_id: str, _: TaskUpdateRequest) -> Task:
|
||||
not_implemented(f"tasks.update:{task_id}")
|
||||
@router.patch("/tasks/{task_id}", response_model=Task, tags=["Tasks"])
|
||||
async def update_task(task_id: str, request: TaskUpdateRequest) -> Task:
|
||||
return task_service.update_task(task_id, request.model_dump(exclude_unset=True))
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/tasks/{task_id}",
|
||||
response_model=OperationResponse,
|
||||
responses=not_implemented_response,
|
||||
tags=["Tasks"],
|
||||
)
|
||||
async def delete_task(task_id: str) -> OperationResponse:
|
||||
not_implemented(f"tasks.delete:{task_id}")
|
||||
if not task_service.delete_task(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
|
||||
@@ -556,21 +575,26 @@ async def delete_task(task_id: str) -> OperationResponse:
|
||||
"/media/transcriptions",
|
||||
response_model=TranscriptionJob,
|
||||
status_code=202,
|
||||
responses=not_implemented_response,
|
||||
tags=["Media"],
|
||||
)
|
||||
async def create_transcription(_: TranscriptionRequest) -> TranscriptionJob:
|
||||
not_implemented("media.transcriptions.create")
|
||||
async def create_transcription(request: TranscriptionRequest) -> TranscriptionJob:
|
||||
return transcription_service.create_transcription(
|
||||
request.attachment_id, request.language
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/media/transcriptions/{job_id}",
|
||||
response_model=TranscriptionJob,
|
||||
responses=not_implemented_response,
|
||||
tags=["Media"],
|
||||
)
|
||||
async def get_transcription(job_id: str) -> TranscriptionJob:
|
||||
not_implemented(f"media.transcriptions.read:{job_id}")
|
||||
job = transcription_service.get_transcription(job_id)
|
||||
if job is None:
|
||||
raise ApiError(
|
||||
404, "RESOURCE_NOT_FOUND", "transcription job not found", {"job_id": job_id}
|
||||
)
|
||||
return job
|
||||
|
||||
|
||||
@router.get("/index/status", response_model=IndexStatus, tags=["Index"])
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
from app.config import get_settings
|
||||
from app.errors import ApiError
|
||||
|
||||
_ATTACHMENT_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]{0,254}$")
|
||||
MAX_ATTACHMENT_BYTES = 1024 * 1024
|
||||
|
||||
|
||||
def attachment_path(attachment_id: str) -> Path:
|
||||
if not _ATTACHMENT_ID.fullmatch(attachment_id):
|
||||
raise ApiError(
|
||||
400, "INVALID_ATTACHMENT_ID", "attachment_id is invalid",
|
||||
{"attachment_id": attachment_id},
|
||||
)
|
||||
root = get_settings().attachments_path.resolve()
|
||||
candidate = (root / attachment_id).resolve()
|
||||
if not candidate.is_relative_to(root):
|
||||
raise ApiError(400, "INVALID_ATTACHMENT_ID", "attachment path escapes storage")
|
||||
return candidate
|
||||
|
||||
|
||||
def read_attachment(attachment_id: str, *, max_chars: int = 100_000) -> dict[str, object]:
|
||||
path = attachment_path(attachment_id)
|
||||
if not path.is_file():
|
||||
raise ApiError(
|
||||
404, "ATTACHMENT_NOT_FOUND", "attachment not found",
|
||||
{"attachment_id": attachment_id},
|
||||
)
|
||||
size = path.stat().st_size
|
||||
if size > MAX_ATTACHMENT_BYTES:
|
||||
raise ApiError(
|
||||
413, "ATTACHMENT_TOO_LARGE", "attachment exceeds the Tool read limit",
|
||||
{"attachment_id": attachment_id, "size": size},
|
||||
)
|
||||
text = path.read_text(encoding="utf-8")
|
||||
truncated = len(text) > max_chars
|
||||
return {
|
||||
"attachment_id": attachment_id,
|
||||
"content": text[:max_chars],
|
||||
"size": size,
|
||||
"truncated": truncated,
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
import asyncio
|
||||
from functools import wraps
|
||||
|
||||
_vault_mutation_lock = asyncio.Lock()
|
||||
|
||||
|
||||
def serialized_vault_mutation(operation):
|
||||
"""串行化 Vault 文件与可重建索引的写入,避免 rebuild 与 Note 写操作交错。"""
|
||||
|
||||
@wraps(operation)
|
||||
async def wrapped(*args, **kwargs):
|
||||
async with _vault_mutation_lock:
|
||||
return await operation(*args, **kwargs)
|
||||
|
||||
return wrapped
|
||||
@@ -17,11 +17,24 @@ 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.coordination import serialized_vault_mutation
|
||||
from app.retrieval.vectorstore import SqliteVecStore
|
||||
|
||||
vector_store = SqliteVecStore()
|
||||
|
||||
_jobs: dict[str, IndexJob] = {}
|
||||
_active_job_id: str | None = None
|
||||
_last_completed_at: datetime | None = None
|
||||
_last_error: str | None = None
|
||||
MAX_JOBS = 100
|
||||
|
||||
|
||||
def _remember_job(job: IndexJob) -> None:
|
||||
_jobs[job.job_id] = job
|
||||
while len(_jobs) > MAX_JOBS:
|
||||
oldest = next(iter(_jobs))
|
||||
_jobs.pop(oldest, None)
|
||||
|
||||
|
||||
def _scan_vault() -> list[tuple[str, str, str, datetime, datetime]]:
|
||||
@@ -29,23 +42,28 @@ def _scan_vault() -> list[tuple[str, str, str, datetime, datetime]]:
|
||||
|
||||
先读入内存:若文件读取失败,rebuild 尚未清空旧索引,不会造成数据损失。
|
||||
"""
|
||||
vault = get_settings().vault_path
|
||||
vault = get_settings().vault_path.resolve()
|
||||
result: list[tuple[str, str, str, datetime, datetime]] = []
|
||||
if not vault.exists():
|
||||
return result
|
||||
for path in sorted(vault.rglob("*.md")):
|
||||
resolved = path.resolve()
|
||||
if not resolved.is_relative_to(vault):
|
||||
continue
|
||||
rel = path.relative_to(vault).as_posix()
|
||||
folder = path.relative_to(vault).parent.as_posix()
|
||||
if folder == ".":
|
||||
folder = ""
|
||||
stat = path.stat()
|
||||
stat = resolved.stat()
|
||||
created = datetime.fromtimestamp(stat.st_ctime, tz=timezone.utc)
|
||||
updated = datetime.fromtimestamp(stat.st_mtime, tz=timezone.utc)
|
||||
result.append((rel, folder, path.read_text(encoding="utf-8"), created, updated))
|
||||
result.append((rel, folder, resolved.read_text(encoding="utf-8"), created, updated))
|
||||
return result
|
||||
|
||||
|
||||
@serialized_vault_mutation
|
||||
async def rebuild(request: IndexRebuildRequest) -> IndexJob:
|
||||
global _active_job_id, _last_completed_at, _last_error
|
||||
job_id = "job_" + uuid4().hex[:12]
|
||||
# 增量重建(scope != all 或指定 note_ids)尚未实现,明确拒绝而非静默全量重建
|
||||
if request.scope != "all" or request.note_ids:
|
||||
@@ -59,10 +77,22 @@ async def rebuild(request: IndexRebuildRequest) -> IndexJob:
|
||||
# 先扫描到内存(失败不会清旧索引),再快照旧库用于失败回滚
|
||||
docs = _scan_vault()
|
||||
settings = get_settings()
|
||||
backup_path = settings.db_path.with_suffix(".db.bak") if settings.db_path.exists() else None
|
||||
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
|
||||
_remember_job(IndexJob(
|
||||
job_id=job_id, status="running", scope=request.scope,
|
||||
created_at=datetime.now(timezone.utc),
|
||||
))
|
||||
try:
|
||||
repository.clear_all()
|
||||
await vector_store.clear()
|
||||
@@ -72,27 +102,39 @@ async def rebuild(request: IndexRebuildRequest) -> IndexJob:
|
||||
created_at=created, updated_at=updated,
|
||||
)
|
||||
await index_note(parsed)
|
||||
except BaseException:
|
||||
task_service.restore_note_links(task_note_links)
|
||||
except BaseException as exc:
|
||||
# 重建失败:恢复旧索引,避免留下半成品;记录 failed 任务后向上抛
|
||||
if backup_path is not None and backup_path.exists():
|
||||
shutil.copy2(backup_path, settings.db_path)
|
||||
_jobs[job_id] = IndexJob(
|
||||
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),
|
||||
)
|
||||
))
|
||||
_last_error = str(exc)
|
||||
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))
|
||||
_jobs[job_id] = job
|
||||
_remember_job(job)
|
||||
_last_completed_at = job.created_at
|
||||
return job
|
||||
|
||||
|
||||
def get_status() -> IndexStatus:
|
||||
# 同步重建、无排队任务,因此状态恒为 idle;实际索引规模可由 GET /api/notes 与搜索反映
|
||||
return IndexStatus(status="idle", pending_jobs=0)
|
||||
if _active_job_id is not None:
|
||||
return IndexStatus(status="running", pending_jobs=0, active_job_id=_active_job_id)
|
||||
return IndexStatus(
|
||||
status="failed" if _last_error else "idle",
|
||||
pending_jobs=0,
|
||||
last_completed_at=_last_completed_at,
|
||||
error_message=_last_error,
|
||||
)
|
||||
|
||||
|
||||
def get_job(job_id: str) -> IndexJob | None:
|
||||
|
||||
@@ -19,6 +19,7 @@ from app.errors import ApiError
|
||||
from app.knowledge.parser import ParsedNote, parse_note
|
||||
from app.retrieval.embedding import HashEmbeddingProvider
|
||||
from app.retrieval.vectorstore import SqliteVecStore, VectorRecord
|
||||
from app.services.coordination import serialized_vault_mutation
|
||||
|
||||
# 轻量实现实例(无状态,可直接复用);接入真实模型后替换为对应 Provider
|
||||
embedding = HashEmbeddingProvider()
|
||||
@@ -148,6 +149,7 @@ async def index_note(parsed: ParsedNote) -> None:
|
||||
conn.close()
|
||||
|
||||
|
||||
@serialized_vault_mutation
|
||||
async def create_note(*, title: str, markdown: str, folder: str | None, tags: list[str]) -> Note:
|
||||
rel_path, clean_folder = _rel_path(folder, title)
|
||||
now = datetime.now(timezone.utc)
|
||||
@@ -183,6 +185,7 @@ async def get_note(note_id: str) -> Note | None:
|
||||
record.created_at, record.updated_at, record.blocks, markdown)
|
||||
|
||||
|
||||
@serialized_vault_mutation
|
||||
async def update_note(
|
||||
note_id: str, *, title: str | None = None, markdown: str | None = None, tags: list[str] | None = None
|
||||
) -> Note:
|
||||
@@ -213,6 +216,58 @@ async def update_note(
|
||||
parsed.created_at, parsed.updated_at, parsed.blocks, new_md)
|
||||
|
||||
|
||||
@serialized_vault_mutation
|
||||
async def move_note(note_id: str, *, folder: str) -> Note:
|
||||
record = repository.get_note_record(note_id)
|
||||
if record is None:
|
||||
raise ApiError(404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id})
|
||||
|
||||
clean_folder = _normalize_folder(folder)
|
||||
filename = Path(record.file_path).name
|
||||
new_rel_path = f"{clean_folder}/{filename}" if clean_folder else filename
|
||||
if new_rel_path == record.file_path:
|
||||
note = await get_note(note_id)
|
||||
assert note is not None
|
||||
return note
|
||||
|
||||
source = _abs_path(record.file_path)
|
||||
target = _abs_path(new_rel_path)
|
||||
if not source.is_file():
|
||||
raise ApiError(
|
||||
409, "NOTE_FILE_MISSING", "note file is missing from the Vault",
|
||||
{"note_id": note_id, "file_path": record.file_path},
|
||||
)
|
||||
if target.exists():
|
||||
raise ApiError(
|
||||
409, "RESOURCE_CONFLICT", "a note already exists at the target path",
|
||||
{"note_id": note_id, "file_path": new_rel_path},
|
||||
)
|
||||
|
||||
markdown = source.read_text(encoding="utf-8")
|
||||
target.parent.mkdir(parents=True, exist_ok=True)
|
||||
source.replace(target)
|
||||
try:
|
||||
parsed = parse_note(
|
||||
markdown=markdown,
|
||||
file_path=new_rel_path,
|
||||
folder=clean_folder,
|
||||
tags=record.tags,
|
||||
created_at=record.created_at,
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
note_id=record.note_id,
|
||||
)
|
||||
parsed.title = record.title
|
||||
await index_note(parsed)
|
||||
except BaseException:
|
||||
target.replace(source)
|
||||
raise
|
||||
return _build_note(
|
||||
parsed.note_id, parsed.title, parsed.file_path, parsed.tags,
|
||||
parsed.created_at, parsed.updated_at, parsed.blocks, markdown,
|
||||
)
|
||||
|
||||
|
||||
@serialized_vault_mutation
|
||||
async def delete_note(note_id: str) -> bool:
|
||||
record = repository.get_note_record(note_id)
|
||||
if record is None:
|
||||
|
||||
@@ -0,0 +1,159 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from uuid import uuid4
|
||||
|
||||
from app import repository
|
||||
from app.contracts import Task, TaskStatus
|
||||
from app.database.db import connect, transaction
|
||||
from app.errors import ApiError
|
||||
|
||||
|
||||
def _now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
def _task_from_row(row) -> Task:
|
||||
return Task(
|
||||
task_id=row["task_id"],
|
||||
title=row["title"],
|
||||
description=row["description"],
|
||||
status=TaskStatus(row["status"]),
|
||||
note_id=row["note_id"],
|
||||
due_at=datetime.fromisoformat(row["due_at"]) if row["due_at"] else None,
|
||||
created_at=datetime.fromisoformat(row["created_at"]),
|
||||
updated_at=datetime.fromisoformat(row["updated_at"]),
|
||||
)
|
||||
|
||||
|
||||
def create_task(
|
||||
*, title: str, description: str = "", note_id: str | None = None,
|
||||
due_at: datetime | None = None,
|
||||
) -> Task:
|
||||
if note_id and repository.get_note_record(note_id) is None:
|
||||
raise ApiError(404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id})
|
||||
task_id = f"task_{uuid4().hex}"
|
||||
now = _now()
|
||||
conn = connect()
|
||||
try:
|
||||
with transaction(conn):
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT INTO tasks
|
||||
(task_id, title, description, status, note_id, due_at, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
task_id, title, description, TaskStatus.todo.value, note_id,
|
||||
due_at.isoformat() if due_at else None, now.isoformat(), now.isoformat(),
|
||||
),
|
||||
)
|
||||
row = conn.execute("SELECT * FROM tasks WHERE task_id = ?", (task_id,)).fetchone()
|
||||
return _task_from_row(row)
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def get_task(task_id: str) -> Task | None:
|
||||
conn = connect()
|
||||
try:
|
||||
row = conn.execute("SELECT * FROM tasks WHERE task_id = ?", (task_id,)).fetchone()
|
||||
return _task_from_row(row) if row else None
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def list_tasks(*, limit: int, offset: int) -> tuple[list[Task], int]:
|
||||
conn = connect()
|
||||
try:
|
||||
total = conn.execute("SELECT COUNT(*) FROM tasks").fetchone()[0]
|
||||
rows = conn.execute(
|
||||
"SELECT * FROM tasks ORDER BY updated_at DESC LIMIT ? OFFSET ?",
|
||||
(limit, offset),
|
||||
).fetchall()
|
||||
return [_task_from_row(row) for row in rows], total
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def update_task(task_id: str, values: dict[str, object]) -> Task:
|
||||
current = get_task(task_id)
|
||||
if current is None:
|
||||
raise ApiError(404, "RESOURCE_NOT_FOUND", "task not found", {"task_id": task_id})
|
||||
if "note_id" in values and values["note_id"]:
|
||||
note_id = str(values["note_id"])
|
||||
if repository.get_note_record(note_id) is None:
|
||||
raise ApiError(404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id})
|
||||
if values.get("title") is None:
|
||||
values.pop("title", None)
|
||||
if values.get("description") is None:
|
||||
values.pop("description", None)
|
||||
if values.get("status") is None:
|
||||
values.pop("status", None)
|
||||
|
||||
columns: list[str] = []
|
||||
params: list[object] = []
|
||||
for name, value in values.items():
|
||||
columns.append(f"{name} = ?")
|
||||
if isinstance(value, datetime):
|
||||
value = value.isoformat()
|
||||
elif isinstance(value, TaskStatus):
|
||||
value = value.value
|
||||
params.append(value)
|
||||
columns.append("updated_at = ?")
|
||||
params.append(_now().isoformat())
|
||||
params.append(task_id)
|
||||
conn = connect()
|
||||
try:
|
||||
with transaction(conn):
|
||||
conn.execute(
|
||||
f"UPDATE tasks SET {', '.join(columns)} WHERE task_id = ?",
|
||||
params,
|
||||
)
|
||||
row = conn.execute("SELECT * FROM tasks WHERE task_id = ?", (task_id,)).fetchone()
|
||||
return _task_from_row(row)
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def delete_task(task_id: str) -> bool:
|
||||
conn = connect()
|
||||
try:
|
||||
with transaction(conn):
|
||||
cursor = conn.execute("DELETE FROM tasks WHERE task_id = ?", (task_id,))
|
||||
return cursor.rowcount > 0
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def note_links() -> dict[str, str]:
|
||||
"""重建可再生 Note 索引前,暂存不可再生 Task 到 Note 的业务关联。"""
|
||||
conn = connect()
|
||||
try:
|
||||
return {
|
||||
row["task_id"]: row["note_id"]
|
||||
for row in conn.execute(
|
||||
"SELECT task_id, note_id FROM tasks WHERE note_id IS NOT NULL"
|
||||
)
|
||||
}
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def restore_note_links(links: dict[str, str]) -> None:
|
||||
if not links:
|
||||
return
|
||||
conn = connect()
|
||||
try:
|
||||
with transaction(conn):
|
||||
for task_id, note_id in links.items():
|
||||
exists = conn.execute(
|
||||
"SELECT 1 FROM notes WHERE note_id = ?", (note_id,)
|
||||
).fetchone()
|
||||
if exists:
|
||||
conn.execute(
|
||||
"UPDATE tasks SET note_id = ? WHERE task_id = ?",
|
||||
(note_id, task_id),
|
||||
)
|
||||
finally:
|
||||
conn.close()
|
||||
@@ -0,0 +1,40 @@
|
||||
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.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:
|
||||
del language # 预生成 transcript 暂不需要语言识别。
|
||||
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."
|
||||
),
|
||||
created_at=datetime.now(timezone.utc),
|
||||
)
|
||||
_jobs[job.job_id] = job
|
||||
while len(_jobs) > MAX_JOBS:
|
||||
_jobs.popitem(last=False)
|
||||
return job.model_copy(deep=True)
|
||||
|
||||
|
||||
def get_transcription(job_id: str) -> TranscriptionJob | None:
|
||||
job = _jobs.get(job_id)
|
||||
return job.model_copy(deep=True) if job else None
|
||||
Reference in New Issue
Block a user