diff --git a/.gitignore b/.gitignore index cc98443..bc4f021 100644 --- a/.gitignore +++ b/.gitignore @@ -8,3 +8,6 @@ data/processed/ data/raw/*.jsonl data/raw/*.db *.log +.coverage +.coverage.* +htmlcov/ diff --git a/INTEGRATION.md b/INTEGRATION.md index 0e42b30..36cc84a 100644 --- a/INTEGRATION.md +++ b/INTEGRATION.md @@ -2,49 +2,115 @@ ## Shared boundary -The extractors consume `core.models.MemoryEvent` and return +The B-side extractors consume `core.models.MemoryEvent` and return `core.models.MemoryCandidate`. They do not alter `core/constants.py`, `core/models.py`, or the Phase 0 SQLite schema. -The B-side extractors are compatible with A-side events produced by -`ingestion.collector.create_raw_event()` followed by -`ingestion.adapter.raw_event_to_memory_event()`. +The A-to-B contract is: -Callers must provide a non-empty `user_id`, `session_id`, and `task_id` for -tenant isolation and concurrent task separation. Timestamps must be UTC or +```text +Raw payload + -> ingestion.collector.create_raw_event(...) + -> ingestion.adapter.raw_event_to_memory_event(...) + -> B-side extractor + -> MemoryCandidate[] +``` + +Callers should provide `user_id`, `session_id`, and `task_id` for tenant +isolation and concurrent task separation. Timestamps must be UTC or timezone-aware values accepted by `MemoryEvent`. -## Extractor semantics +## B-side extractor semantics + +The public class and method names remain compatible with the Phase 1 task +schedule: + +- `PreferenceExtractor` +- `KnowledgeExtractor` +- `WorkflowExtractor` +- `ToolExtractor` +- `EnvironmentExtractor` + +The implementation is now LLM-only for memory extraction. These classes do not +use regex templates, keyword lists, frequency counters, event-type shortcuts, +text-length thresholds, or locally computed confidence gates to decide what +memory should be extracted. + +All semantic extraction is delegated to a configured JSON-capable LLM adapter: + +```python +from extractors.llm_memory_extractor import set_default_llm_client + +set_default_llm_client(client) # client.complete_json(prompt, schema) -> JSON +``` + +If no LLM client is configured, legacy static methods return empty results +rather than silently falling back to hardcoded extraction rules. This keeps +failure behavior explicit for integration and evaluation. + +## Pipeline + +```text +MemoryEvent / MemoryEvent[] + | + v +Prompt payload packaging and secret/contact redaction + | + v +LLMJsonClient.complete_json(prompt, schema) + | + v +LLM structured JSON + | + v +CandidateValidator + | + v +CandidateMerger + | + v +MemoryCandidate[] +``` -- `PreferenceExtractor` and `KnowledgeExtractor` emit user-scoped candidates. - Template evidence is never aggregated across users. -- `WorkflowExtractor` groups by `(user_id, session_id, task_id)` and orders - events by `(timestamp, event_id)` before deriving a workflow. A workflow is - emitted only when its evidence-based `metadata.reproduction_rate` is at - least `0.8`. Repeated identical workflows merge evidence into - `metadata.occurrence_count` rather than dropping later observations. -- `ToolExtractor.calculate_tool_success_rate()` uses - `success_count / completed_invocation_count`. Calls without terminal results - are exposed as `unknown_count` and excluded from the denominator, preventing - incomplete telemetry from being reported as failures. -- `EnvironmentExtractor.extract_from_tool_output()` accepts a dictionary. - It supports common keys such as `downloads`, `documents`, `locale`, - `installed_software`, `applications`, and `os_version`. Home-directory - usernames in paths are normalised to `~`; credential-like keys are ignored. +Local post-processing is limited to engineering boundaries: -`MemoryCandidate.key` and `candidate_id` are deterministic for an identical -user/type/key input. The downstream storage owner remains responsible for -upserting records by the candidate key or another approved identity rule. +- prompt packaging from `MemoryEvent`; +- credential/contact redaction before prompt and after model output; +- JSON/schema conversion into `MemoryCandidate`; +- empty/non-long-term candidate filtering based on model output flags; +- duplicate merging and conflict annotation for already-produced candidates. -## Verification +It is not a rule-based memory extractor. -Runtime code uses only the standard library. For development tests: +## A-side integration verification + +The test `tests/test_ingestion_to_extractors.py` verifies that ingestion-style +`MemoryEvent` objects can directly drive B-side LLM extractors for: + +- preference extraction from conversation events; +- knowledge/tool extraction from tool result events; +- environment extraction from metadata/tool-output dictionaries. + +The test uses a fake LLM client, so it is deterministic and does not require a +network connection, API key, or model download. + +## Verification commands + +Targeted B-side verification: + +```powershell +python -m pytest tests/test_extractors.py tests/test_llm_memory_extractor.py tests/test_environment_extractor.py tests/test_tool_extractor.py tests/test_workflow_extractor.py tests/test_ingestion_to_extractors.py -q +``` + +Coverage: ```powershell -python -m pip install -r requirements-dev.txt -python -m pytest tests -v -python -m coverage run -m pytest tests -q -python -m coverage report -m +python -m coverage run -m pytest tests/test_extractors.py tests/test_llm_memory_extractor.py tests/test_environment_extractor.py tests/test_tool_extractor.py tests/test_workflow_extractor.py tests/test_ingestion_to_extractors.py -q +python -m coverage report -m --include="extractors/*" ``` -Validated locally with Python 3.13.5, pytest 8.3.4, and coverage 7.14.3. +Dataset input/output demo: + +```powershell +python demo\b_llm_pipeline_demo.py +``` diff --git a/demo/b_llm_pipeline_demo.py b/demo/b_llm_pipeline_demo.py new file mode 100644 index 0000000..6445b72 --- /dev/null +++ b/demo/b_llm_pipeline_demo.py @@ -0,0 +1,252 @@ +"""Print B-side LLM-only memory extraction input/output. + +Run: + python demo/b_llm_pipeline_demo.py + +The demo uses a fake LLM client. It does not call any external API. +It is designed for screen sharing: it shows the dataset before LLM processing, +the sanitized prompt sent to the model, the model JSON response, and the final +``MemoryCandidate[]`` after validation/conversion. +""" + +from __future__ import annotations + +import json +import sys +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +if str(PROJECT_ROOT) not in sys.path: + sys.path.insert(0, str(PROJECT_ROOT)) + +from core.constants import EventType, Scene +from core.models import MemoryEvent +from extractors.knowledge_extractor import KnowledgeExtractor +from extractors.llm_memory_extractor import ( + HybridMemoryExtractor, + LLMMemoryExtractor, + build_memory_extraction_prompt, + clear_default_llm_client, + event_payload_for_demo, + set_default_llm_client, +) +from extractors.preference_extractor import PreferenceExtractor +from extractors.workflow_extractor import WorkflowExtractor + + +class FakeLLMClient: + def __init__(self) -> None: + self.calls: list[dict[str, Any]] = [] + + def complete_json(self, prompt: str, schema: dict[str, Any]) -> dict[str, Any]: + request = json.loads(prompt) + self.calls.append({"request": request, "schema": schema}) + mode = request["mode"] + if mode.startswith("preference"): + return { + "candidates": [ + _candidate( + "preference", + "response_order", + "conclusion_first", + "用户偏好报告类内容先给结论再展开", + "这种报告以后别写太散,先给结论再展开。", + "LLM 判断这是长期写作偏好。", + 0.88, + ) + ] + } + if mode.startswith("knowledge"): + return { + "candidates": [ + _candidate( + "knowledge", + "tool_case", + "batch_export", + "批量导出时可复用:输入文件列表,输出合并后的 report.pdf", + "工具结果显示已完成 batch export。", + "LLM 从工具输入输出中总结可复用知识。", + 0.84, + ) + ] + } + if mode.startswith("workflow") or mode == "session": + return { + "candidates": [ + _candidate( + "workflow", + "report_generation", + "conclusion_then_export", + "报告生成流程:先给结论,再展开依据,最后导出 PDF", + "对话偏好和工具导出结果共同构成流程。", + "LLM 将多条事件归纳为可复用流程。", + 0.86, + ) + ] + } + return { + "candidates": [ + _candidate( + "preference", + "response_order", + "conclusion_first", + "用户偏好报告类内容先给结论再展开", + "这种报告以后别写太散,先给结论再展开。", + "LLM 判断这是长期写作偏好。", + 0.88, + ), + _candidate( + "knowledge", + "tool_case", + "batch_export", + "批量导出时可复用:输入文件列表,输出合并后的 report.pdf", + "工具结果显示已完成 batch export。", + "LLM 从工具输入输出中总结可复用知识。", + 0.84, + ), + _candidate( + "workflow", + "report_generation", + "conclusion_then_export", + "报告生成流程:先给结论,再展开依据,最后导出 PDF", + "对话偏好和工具导出结果共同构成流程。", + "LLM 将多条事件归纳为可复用流程。", + 0.86, + ), + { + **_candidate( + "knowledge", + "credential", + "api_key", + "用户 API key 是 sk-demo-secret", + "api_key=sk-demo-secret", + "这是敏感凭据,应被校验层拒绝。", + 0.91, + ), + "sensitivity": "api_key", + }, + ] + } + + +def _candidate( + memory_type: str, + category: str, + value: str, + content: str, + evidence: str, + reason: str, + confidence: float, +) -> dict[str, Any]: + return { + "is_memory_worthy": True, + "is_long_term": True, + "memory_type": memory_type, + "category": category, + "value": value, + "scope": "demo", + "content": content, + "confidence": confidence, + "evidence": evidence, + "reason": reason, + "sensitivity": "none", + } + + +def _event( + event_id: str, + event_type: EventType, + *, + content: str | None = None, + tool_name: str | None = None, + input_payload: dict[str, Any] | None = None, + output_payload: dict[str, Any] | None = None, + metadata: dict[str, Any] | None = None, +) -> MemoryEvent: + return MemoryEvent( + event_id=event_id, + raw_event_id=f"raw-{event_id}", + user_id="user-demo", + session_id="session-demo", + task_id="task-demo", + event_type=event_type, + scenario=Scene.OFFICE, + source=event_type.value, + actor="user" if event_type is EventType.CONVERSATION else "tool", + content=content, + tool_name=tool_name, + input=input_payload or {}, + output=output_payload or {}, + metadata=metadata or {}, + success=True if event_type is EventType.TOOL_RESULT else None, + timestamp=datetime(2026, 7, 5, 19, 0, tzinfo=timezone.utc), + ) + + +def _print_json(title: str, value: Any) -> None: + print(f"\n=== {title} ===") + print(json.dumps(value, ensure_ascii=False, indent=2, default=str)) + + +def main() -> None: + events = [ + _event( + "demo-conv-001", + EventType.CONVERSATION, + content="这种报告以后别写太散,先给结论再展开。", + ), + _event( + "demo-tool-001", + EventType.TOOL_RESULT, + tool_name="batch_export", + input_payload={"files": ["a.docx", "b.docx"], "format": "pdf"}, + output_payload={ + "status": "success", + "file": "report.pdf", + "api_key": "sk-demo-secret", + }, + ), + ] + client = FakeLLMClient() + + prompt = build_memory_extraction_prompt(events, mode="demo_dataset") + raw_llm_output = client.complete_json(prompt, {}) + final_candidates = LLMMemoryExtractor.extract_events(events, client, mode="demo_dataset") + + set_default_llm_client(client) + try: + legacy_api_outputs = { + "PreferenceExtractor.extract_from_conversation": [ + candidate.to_dict() for candidate in PreferenceExtractor.extract_from_conversation(events[0]) + ], + "KnowledgeExtractor.extract_from_tool_result": [ + candidate.to_dict() for candidate in KnowledgeExtractor.extract_from_tool_result(events[1]) + ], + "WorkflowExtractor.extract_multi_step_workflow": [ + candidate.to_dict() for candidate in WorkflowExtractor.extract_multi_step_workflow(events) + ], + "HybridMemoryExtractor.extract_from_session": [ + candidate.to_dict() for candidate in HybridMemoryExtractor.extract_from_session(events) + ], + } + finally: + clear_default_llm_client() + + _print_json("1. LLM 处理前:原始 MemoryEvent 数据集", [event.to_dict() for event in events]) + _print_json("2. LLM 处理前:实际入模的脱敏 JSON", [event_payload_for_demo(event) for event in events]) + _print_json("3. LLM 请求 prompt 结构", json.loads(prompt)) + _print_json("4. LLM 原始输出 JSON", raw_llm_output) + _print_json("5. LLM 输出后:校验/脱敏/转换后的 MemoryCandidate[]", [candidate.to_dict() for candidate in final_candidates]) + _print_json("6. 旧 B 侧 API 现在的 LLM-only 输出", legacy_api_outputs) + + print("\n=== 7. 展示结论 ===") + print("- B 侧旧函数名保留,但内部不再跑正则、关键词、频次统计或长度阈值。") + print("- LLM 处理前数据是 MemoryEvent / MemoryEvent[],入模前会脱敏。") + print("- LLM 输出必须是结构化 JSON,再转换为 MemoryCandidate[]。") + print("- 敏感凭据候选会在校验层拒绝,因此不会进入最终 MemoryCandidate[]。") + + +if __name__ == "__main__": + main() diff --git a/demo/b_memory_admission_demo.py b/demo/b_memory_admission_demo.py new file mode 100644 index 0000000..3859182 --- /dev/null +++ b/demo/b_memory_admission_demo.py @@ -0,0 +1,135 @@ +"""Screen-share demo for candidate admission and conflict handling. + +Run: + python demo/b_memory_admission_demo.py + +The demo uses a fake JSON LLM and a temporary SQLite database. It makes no +external request and leaves no database file behind. +""" + +from __future__ import annotations + +import json +import sys +import tempfile +from pathlib import Path +from typing import Any + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +if str(PROJECT_ROOT) not in sys.path: + sys.path.insert(0, str(PROJECT_ROOT)) + +from core.constants import MemoryStatus, MemoryType, Scene +from core.models import MemoryCandidate +from memory.admission import MemoryAdmissionService, SQLiteAdmissionRepository +from memory.conflict_resolver import LLMConflictResolver + + +class DemoAdmissionLLM: + """Deterministic fake used only to display the production LLM contract.""" + + def __init__(self) -> None: + self.calls: list[dict[str, Any]] = [] + + def complete_json(self, prompt: str, schema: dict[str, Any]) -> dict[str, Any]: + request = json.loads(prompt) + candidate = request["candidate"] + active = request["active_memories"] + if not active: + output = { + "action": "create", + "reason": "没有同一记忆键的已生效记录,可以首次入库", + "target_memory_ids": [], + "final_content": candidate["content"], + "requires_human_review": False, + "decision_confidence": 0.96, + } + else: + output = { + "action": "replace", + "reason": "新偏好明确更新了同一办公导出设置", + "target_memory_ids": [active[0]["memory_id"]], + "final_content": candidate["content"], + "requires_human_review": False, + "decision_confidence": 0.94, + } + self.calls.append({"request": request, "response": output, "schema": schema}) + return output + + +def make_candidate(candidate_id: str, content: str) -> MemoryCandidate: + return MemoryCandidate( + candidate_id=candidate_id, + user_id="user-demo", + memory_type=MemoryType.PREFERENCE, + key="preference.export.format", + content=content, + scenario=Scene.OFFICE, + confidence=0.91, + source="llm_extracted", + source_events=[f"event-{candidate_id}"], + tags=["preference", "export"], + metadata={"evidence": content}, + ) + + +def print_json(title: str, value: Any) -> None: + print(f"\n=== {title} ===") + print(json.dumps(value, ensure_ascii=False, indent=2, default=str)) + + +def main() -> None: + fake_llm = DemoAdmissionLLM() + with tempfile.TemporaryDirectory(prefix="os-memory-admission-") as temp_dir: + repository = SQLiteAdmissionRepository(str(Path(temp_dir) / "demo.db")) + admission = MemoryAdmissionService( + repository, + LLMConflictResolver(fake_llm), + ) + + original = make_candidate("cand-pdf", "用户偏好办公文件使用 PDF 导出") + print_json("1. 第一次入库前:MemoryCandidate", original.to_dict()) + created = admission.admit(original) + print_json("2. 第一次入库:发送给 LLM 的数据", fake_llm.calls[-1]["request"]) + print_json("3. 第一次入库:LLM 原始结构化输出", fake_llm.calls[-1]["response"]) + print_json("4. 第一次入库后:AdmissionResult", created.to_dict()) + + changed = make_candidate("cand-word", "用户现在偏好办公文件使用 Word 导出") + print_json("5. 冲突候选:新的 MemoryCandidate", changed.to_dict()) + replaced = admission.admit(changed) + print_json("6. 冲突处理:发送给 LLM 的候选和已有记忆", fake_llm.calls[-1]["request"]) + print_json("7. 冲突处理:LLM 原始结构化输出", fake_llm.calls[-1]["response"]) + print_json("8. 冲突处理后:AdmissionResult", replaced.to_dict()) + print_json( + "9. SQLite 当前全部版本", + [record.to_dict() for record in repository.list_all(user_id="user-demo")], + ) + + calls_before_retry = len(fake_llm.calls) + retried = admission.admit(changed) + print_json("10. 同一候选重试:幂等结果", retried.to_dict()) + print( + f"LLM 调用次数变化:{calls_before_retry} -> {len(fake_llm.calls)} " + "(重复请求未再次调用模型)" + ) + + archived = admission.transition(replaced.record.memory_id, MemoryStatus.ARCHIVED) + restored = admission.transition(archived.memory_id, MemoryStatus.ACTIVE) + print_json( + "11. 生命周期演示:ACTIVE -> ARCHIVED -> ACTIVE", + { + "archived_status": archived.status.value, + "restored_status": restored.status.value, + "memory_id": restored.memory_id, + }, + ) + + print("\n=== 12. 展示结论 ===") + print("- LLM 决定语义关系:创建、重复、合并、替换、并存、待确认或拒绝。") + print("- 程序负责结构校验、敏感信息保护、版本号、状态机和 SQLite 事务。") + print("- 新旧偏好冲突时,新记录变成 version=2,旧记录变成 superseded。") + print("- 相同候选重复提交不会重复入库,也不会重复调用 LLM。") + + +if __name__ == "__main__": + main() diff --git a/docs/B_LLM_DATASET_IO_CN.md b/docs/B_LLM_DATASET_IO_CN.md new file mode 100644 index 0000000..779a366 --- /dev/null +++ b/docs/B_LLM_DATASET_IO_CN.md @@ -0,0 +1,231 @@ +# B 组 LLM 处理前后数据集输入输出说明 + +这份文档用于展示 B 侧记忆抽取在 LLM 化之后,数据从进入抽取器到输出 +`MemoryCandidate` 的完整形态。核心结论是:代码不再用硬编码规则决定“提取什么”, +只负责数据整理、脱敏、结构校验、去重和冲突标记。 + +## 1. 展示命令 + +在项目根目录运行: + +```bash +python demo/b_llm_pipeline_demo.py +``` + +这个 demo 使用 fake LLM client,不需要真实 API key,也不会访问外部网络。 + +## 2. 数据流总览 + +```text +原始数据集 MemoryEvent[] + | + v +入模前 JSON 整理与脱敏 + | + v +LLM prompt + schema + | + v +LLM 原始 JSON 输出 + | + v +本地校验、脱敏、去重、冲突标记 + | + v +最终 MemoryCandidate[] +``` + +## 3. LLM 处理前:原始数据集 + +输入来自 A 侧或 ingestion 侧标准化后的 `MemoryEvent`。 + +示例一:用户对话事件 + +```json +{ + "event_id": "demo-conv-001", + "user_id": "user-demo", + "session_id": "session-demo", + "task_id": "task-demo", + "event_type": "conversation", + "scenario": "office", + "source": "conversation", + "actor": "user", + "content": "这种报告以后别写太散,先给结论再展开。" +} +``` + +示例二:工具结果事件 + +```json +{ + "event_id": "demo-tool-001", + "user_id": "user-demo", + "session_id": "session-demo", + "task_id": "task-demo", + "event_type": "tool_result", + "tool_name": "batch_export", + "input": { + "files": ["a.docx", "b.docx"], + "format": "pdf" + }, + "output": { + "status": "success", + "file": "report.pdf", + "api_key": "sk-demo-secret" + }, + "success": true +} +``` + +说明: + +- 这一步展示的是原始事件对象; +- `api_key` 这种字段可能存在于工具输出里; +- 这里还没有做记忆判断。 + +## 4. LLM 处理前:实际入模 JSON + +进入 LLM 前,代码只做结构整理和安全脱敏,不做“关键词命中”“长度大于多少” +这类抽取判断。 + +工具输出中的敏感字段会变成: + +```json +{ + "output": { + "status": "success", + "file": "report.pdf", + "api_key": "[REDACTED_SECRET]" + } +} +``` + +这一步可以向对方说明: + +- 脱敏是安全边界,不是记忆抽取规则; +- 模型看到的是可解释、可复用的事件 JSON; +- 密钥、手机号、邮箱等信息不会原样进入 prompt。 + +## 5. LLM 请求:prompt 和输出契约 + +Prompt 中明确要求模型做语义判断: + +```json +{ + "extraction_policy": [ + "Use semantic understanding to decide whether information is long-term memory.", + "Do not depend on keyword lists, regex templates, text length thresholds, or rule candidates.", + "Extract preference, knowledge, workflow, template, tool, environment, safety, profile, task_state, or session_summary memory when supported by evidence.", + "Mark temporary/current-task-only information as is_memory_worthy=false or is_long_term=false.", + "Never output credentials or raw secrets. Use redacted evidence when needed.", + "Keep evidence and reason explicit so reviewers can understand why the memory exists." + ] +} +``` + +这一步对应主管关心的点:不是用 `len(text) > 120`、`confidence >= 0.9` +之类阈值决定,而是让 LLM 按语义和证据输出结构化结果。 + +## 6. LLM 原始输出 + +LLM 返回 JSON,形态如下: + +```json +{ + "candidates": [ + { + "is_memory_worthy": true, + "is_long_term": true, + "memory_type": "preference", + "category": "response_order", + "value": "conclusion_first", + "scope": "demo", + "content": "用户偏好报告类内容先给结论再展开", + "confidence": 0.88, + "evidence": "这种报告以后别写太散,先给结论再展开。", + "reason": "LLM 判断这是长期写作偏好。", + "sensitivity": "none" + } + ] +} +``` + +LLM 原始输出里必须包含: + +- `memory_type`:记忆类型; +- `content`:可复用记忆内容; +- `confidence`:模型自评置信度; +- `evidence`:来自输入事件的证据; +- `reason`:为什么这是长期记忆; +- `is_memory_worthy` / `is_long_term`:是否值得进入长期记忆。 + +## 7. LLM 处理后:最终 MemoryCandidate + +代码把 LLM JSON 转成团队统一的 `MemoryCandidate`: + +```json +{ + "candidate_id": "6d426b361f1d0f4c", + "user_id": "user-demo", + "memory_type": "preference", + "key": "preference.response_order.conclusion_first", + "content": "用户偏好报告类内容先给结论再展开", + "scenario": "office", + "confidence": 0.88, + "source": "llm_extracted", + "source_events": ["demo-conv-001", "demo-tool-001"], + "source_summaries": ["这种报告以后别写太散,先给结论再展开。"], + "tags": ["llm", "preference", "response_order"], + "metadata": { + "extraction_method": "llm_semantic", + "schema_version": 2, + "evidence": "这种报告以后别写太散,先给结论再展开。", + "reason": "LLM 判断这是长期写作偏好。", + "sensitivity": "none" + } +} +``` + +本地后处理负责: + +- JSON 解析失败保护; +- 空内容过滤; +- 非长期记忆过滤; +- 敏感凭据候选拒绝; +- 输出字段补齐; +- 相同 key/content 去重; +- 相同 key 不同内容时标记冲突。 + +注意:这些是数据质量和安全处理,不是用规则替代 LLM 抽取。 + +## 8. 当前 demo 能展示的三类输出 + +| 输入证据 | LLM 判断 | 最终输出 | +| --- | --- | --- | +| “这种报告以后别写太散,先给结论再展开。” | 长期表达风格偏好 | `preference.response_order.conclusion_first` | +| `batch_export` 工具输入输出 | 可复用工具使用案例 | `knowledge.tool_case.batch_export` | +| 对话偏好 + 工具导出结果 | 可复用报告生成流程 | `workflow.report_generation.conclusion_then_export` | + +同时 demo 还故意模拟一个错误候选: + +```json +{ + "memory_type": "knowledge", + "category": "credential", + "content": "用户 API key 是 sk-demo-secret", + "sensitivity": "api_key" +} +``` + +这个候选会被本地校验层拒绝,不会进入最终 `MemoryCandidate[]`。 + +## 9. 对外说明话术 + +可以这样讲: + +> 我这里展示的是 B 侧 LLM 处理前后的数据形态。处理前是标准 +> MemoryEvent,包括用户对话和工具结果;入模前会先做脱敏,保证敏感字段不会原样进模型; +> LLM 负责判断哪些内容能成为长期记忆,并返回结构化 JSON;代码再把这个 JSON 校验、 +> 去重、补齐来源,转换成 MemoryCandidate。现在抽取判断不再依赖正则、关键词、频次统计 +> 或长度阈值,代码只保留安全和工程边界。 diff --git a/docs/B_LLM_MEMORY_PIPELINE.md b/docs/B_LLM_MEMORY_PIPELINE.md new file mode 100644 index 0000000..ee4bd6f --- /dev/null +++ b/docs/B_LLM_MEMORY_PIPELINE.md @@ -0,0 +1,163 @@ +# B-side LLM-only memory extraction pipeline + +## Current direction + +B-side semantic memory extraction has moved from: + +```text +rule-oriented extraction +``` + +to: + +```text +LLM as the only semantic extraction path +``` + +The old public class names and method signatures are kept for Phase 1 +compatibility, but their internal extraction logic no longer depends on +hardcoded regexes, keyword lists, frequency rules, event-type shortcuts, +text-length thresholds, or local confidence gates. + +## External design ideas mapped to this project + +| Reference idea | Adopted project design | +| --- | --- | +| LangMem: LLM-assisted long-term memory extraction and consolidation | LLM receives `MemoryEvent[]` and returns structured candidate memory | +| Mem0: user-scoped long-term preference and personalization memory | candidates preserve `user_id`, `session_id`, `task_id`, scope, confidence, evidence, and reason | +| Graphiti: temporal/provenance-aware evolving facts | candidates preserve source event ids, event timestamps, evidence, and possible conflict metadata | + +These are design references only. The local code does not import a specific +model SDK and does not call a real model by itself. + +## Architecture + +```text +MemoryEvent / MemoryEvent[] + | + v +Sanitize and package event payloads + | + v +LLMJsonClient.complete_json(prompt, schema) + | + v +LLM structured JSON + | + v +CandidateValidator + | + v +CandidateMerger + | + v +MemoryCandidate[] +``` + +## Local modules + +- `extractors/llm_memory_extractor.py` + - `LLMJsonClient`: minimal protocol for a JSON-capable model adapter. + - `set_default_llm_client(...)`: process-local adapter registration used by + legacy static extractor signatures. + - `LLMMemoryExtractor`: packages events, calls the injected LLM client, and + converts model JSON into `MemoryCandidate`. + - `LLMWorkflowBoundaryExtractor`: asks the model for workflow boundaries. + - `CandidateValidator`: validates model output, rejects credential-like + non-safety candidates, redacts common sensitive values, and filters + candidates explicitly marked non-long-term by the model. + - `CandidateMerger`: de-duplicates candidates and annotates possible + conflicts. + +- Thin compatibility facades: + - `extractors/preference_extractor.py` + - `extractors/knowledge_extractor.py` + - `extractors/workflow_extractor.py` + - `extractors/tool_extractor.py` + - `extractors/environment_extractor.py` + +## Compatibility with B task signatures + +The following methods remain available: + +- `PreferenceExtractor.extract_from_conversation(...)` +- `PreferenceExtractor.extract_from_tool_result(...)` +- `PreferenceExtractor.extract_explicit_preference(...)` +- `PreferenceExtractor.extract_implicit_preference(...)` +- `KnowledgeExtractor.extract_from_tool_result(...)` +- `KnowledgeExtractor.extract_from_conversation(...)` +- `KnowledgeExtractor.extract_templates(...)` +- `WorkflowExtractor.extract_tool_sequence(...)` +- `WorkflowExtractor.extract_multi_step_workflow(...)` +- `WorkflowExtractor.detect_workflow_boundary(...)` +- `ToolExtractor.calculate_tool_success_rate(...)` +- `ToolExtractor.extract_tool_pattern(...)` +- `EnvironmentExtractor.extract_from_tool_output(...)` + +If no default LLM client is configured, these methods return empty results +instead of falling back to hardcoded extraction. + +## Prompt/output contract + +The LLM must return JSON shaped like: + +```json +{ + "candidates": [ + { + "is_memory_worthy": true, + "is_long_term": true, + "memory_type": "preference", + "category": "response_order", + "value": "conclusion_first", + "scope": "report_writing", + "content": "用户偏好报告类内容先给结论再展开", + "confidence": 0.84, + "evidence": "这种报告还是先给结论好一点", + "reason": "用户表达了可复用的后续报告写作偏好", + "sensitivity": "none" + } + ] +} +``` + +The local pipeline converts this into normal `MemoryCandidate` objects. + +## What was intentionally removed + +- preference regex rules; +- knowledge/FAQ keyword extraction; +- repeated-task frequency extraction; +- workflow keyword boundary detection; +- path/category environment heuristics; +- event-type shortcuts for whether to call an LLM; +- `len(text) > 120`, `len(text) > 240`, `confidence >= 0.9`, and similar + unexplained thresholds. + +Remaining regexes are limited to secret/contact redaction and slug/id +normalization. Remaining numeric logic is limited to structural validation, +confidence clamping to `[0, 1]`, and boundary index checking. + +## Verification + +Targeted tests: + +```powershell +python -m pytest tests/test_extractors.py tests/test_llm_memory_extractor.py tests/test_environment_extractor.py tests/test_tool_extractor.py tests/test_workflow_extractor.py tests/test_ingestion_to_extractors.py -q +``` + +Coverage: + +```powershell +python -m coverage run -m pytest tests/test_extractors.py tests/test_llm_memory_extractor.py tests/test_environment_extractor.py tests/test_tool_extractor.py tests/test_workflow_extractor.py tests/test_ingestion_to_extractors.py -q +python -m coverage report -m --include="extractors/*" +``` + +Input/output demo: + +```powershell +python demo\b_llm_pipeline_demo.py +``` + +The demo prints raw `MemoryEvent[]`, sanitized prompt JSON, fake LLM JSON, and +final `MemoryCandidate[]`. diff --git a/docs/B_LLM_MEMORY_PIPELINE_CN.md b/docs/B_LLM_MEMORY_PIPELINE_CN.md new file mode 100644 index 0000000..82ef2f2 --- /dev/null +++ b/docs/B_LLM_MEMORY_PIPELINE_CN.md @@ -0,0 +1,271 @@ +# B 组 LLM-only 记忆抽取方案 + +## 1. 当前改动结论 + +B 侧抽取逻辑已经从: + +```text +规则 baseline + LLM 语义增强 +``` + +调整为: + +```text +LLM 作为唯一正式语义抽取路径 +``` + +也就是说,`PreferenceExtractor`、`KnowledgeExtractor`、`WorkflowExtractor`、 +`ToolExtractor`、`EnvironmentExtractor` 这些旧类名仍然保留,但它们内部不再运行 +正则、关键词、频次统计、工具序列规则或路径分类规则。 + +保留旧类名和旧函数签名的原因是向后兼容团队接口,避免影响 A 侧 `MemoryEvent` +输入和后续存储/检索模块。 + +## 2. 为什么要这样改 + +会议反馈的核心是:偏好、知识、工作流这类记忆抽取不能长期依赖硬编码。 + +硬编码的问题是: + +- 自然语言表达变化太多,正则很难覆盖; +- “以后”“默认”“别写太散”这类表达需要语义判断; +- 工具结果、流程模板、长期偏好之间经常混在一起,需要模型归纳; +- 规则里出现 `120`、`240`、`0.9` 这类阈值时,很难解释来源; +- 后期要做泛化评测,规则堆叠会越来越难维护。 + +因此现在的方案是让 LLM 做判断,代码只负责数据协议和安全边界。 + +## 3. 新架构 + +```text +MemoryEvent / MemoryEvent[] + | + v +Sanitize before prompt + | + v +LLMJsonClient.complete_json(prompt, schema) + | + v +LLM structured JSON + | + v +CandidateValidator + | + v +CandidateMerger + | + v +MemoryCandidate[] +``` + +其中: + +- `MemoryEvent` 是 A 侧/ingestion 侧传进来的标准事件; +- `LLMJsonClient` 是模型适配接口,可以接云端模型、本地模型或 fake client; +- `CandidateValidator` 只做安全和结构校验,不做规则抽取; +- `CandidateMerger` 只做去重和冲突标记,不做规则增强; +- 输出仍然是 `MemoryCandidate[]`,方便后续存储模块继续对接。 + +## 4. 保留了什么,删除了什么 + +保留: + +- 旧函数名; +- 旧返回类型; +- `MemoryEvent` / `MemoryCandidate` 数据结构; +- Phase 0 表结构; +- 敏感信息过滤; +- 候选去重; +- 冲突标记; +- fake LLM 测试方式。 + +删除或停用: + +- 偏好正则; +- FAQ 正则; +- 工具成功率本地统计; +- 工作流关键词边界检测; +- 路径类型硬编码识别; +- `should_call_llm` 里的事件类型捷径; +- `len(text) > 120`、`len(text) > 240`、`confidence >= 0.9` 这类门槛。 + +## 5. 旧 API 现在怎么工作 + +### PreferenceExtractor + +```python +PreferenceExtractor.extract_from_conversation(event) +PreferenceExtractor.extract_from_tool_result(event) +PreferenceExtractor.extract_explicit_preference(content) +PreferenceExtractor.extract_implicit_preference(events) +``` + +现在全部调用默认 LLM client,只返回 `MemoryType.PREFERENCE`。 + +### KnowledgeExtractor + +```python +KnowledgeExtractor.extract_from_tool_result(event) +KnowledgeExtractor.extract_from_conversation(event) +KnowledgeExtractor.extract_templates(events) +``` + +现在全部调用 LLM,由模型判断是否是知识或模板。 + +### WorkflowExtractor + +```python +WorkflowExtractor.detect_workflow_boundary(events) +WorkflowExtractor.extract_tool_sequence(events) +WorkflowExtractor.extract_multi_step_workflow(events) +``` + +边界和流程都由 LLM 输出,不再用关键词或工具序列规则。 + +### ToolExtractor + +```python +ToolExtractor.extract_tool_pattern(events) +ToolExtractor.calculate_tool_success_rate(tool_name, events) +``` + +工具经验由 LLM 输出。成功率不再本地统计,而是读取 LLM 输出候选中的 +`metadata.success_rate`。 + +### EnvironmentExtractor + +```python +EnvironmentExtractor.extract_from_tool_output(output) +``` + +环境信息由 LLM 从系统上下文中判断,不再本地根据 key/path 推断。 + +## 6. LLM 处理前后数据结构 + +### 6.1 LLM 处理前输入 + +输入是 `MemoryEvent` 或 `MemoryEvent[]`。 + +示例: + +```json +{ + "event_id": "demo-conv-001", + "user_id": "user-demo", + "event_type": "conversation", + "content": "这种报告以后别写太散,先给结论再展开。" +} +``` + +工具结果示例: + +```json +{ + "event_id": "demo-tool-001", + "event_type": "tool_result", + "tool_name": "batch_export", + "input": { + "files": ["a.docx", "b.docx"], + "format": "pdf" + }, + "output": { + "status": "success", + "file": "report.pdf", + "api_key": "sk-demo-secret" + } +} +``` + +### 6.2 入模前脱敏 + +进入 LLM 前会把敏感字段替换掉: + +```json +{ + "output": { + "status": "success", + "file": "report.pdf", + "api_key": "[REDACTED_SECRET]" + } +} +``` + +这一步是安全处理,不是记忆抽取规则。 + +### 6.3 LLM 原始输出 + +LLM 必须返回结构化 JSON: + +```json +{ + "candidates": [ + { + "is_memory_worthy": true, + "is_long_term": true, + "memory_type": "preference", + "category": "response_order", + "value": "conclusion_first", + "scope": "demo", + "content": "用户偏好报告类内容先给结论再展开", + "confidence": 0.88, + "evidence": "这种报告以后别写太散,先给结论再展开。", + "reason": "LLM 判断这是长期写作偏好。", + "sensitivity": "none" + } + ] +} +``` + +### 6.4 LLM 处理后输出 + +代码把 LLM JSON 转为 `MemoryCandidate`: + +```json +{ + "memory_type": "preference", + "key": "preference.response_order.conclusion_first", + "content": "用户偏好报告类内容先给结论再展开", + "confidence": 0.88, + "source": "llm_extracted", + "metadata": { + "extraction_method": "llm_semantic", + "evidence": "这种报告以后别写太散,先给结论再展开。", + "reason": "LLM 判断这是长期写作偏好。" + } +} +``` + +## 7. 本地演示脚本 + +运行: + +```bash +python demo/b_llm_pipeline_demo.py +``` + +它会打印: + +1. LLM 处理前的原始 `MemoryEvent` 数据集; +2. 入模前脱敏后的 JSON; +3. LLM prompt 结构; +4. fake LLM 原始输出 JSON; +5. 校验/脱敏/转换后的 `MemoryCandidate[]`; +6. 旧 B 侧 API 现在的 LLM-only 输出。 + +## 8. 当前测试命令 + +```bash +python -m pytest tests/test_extractors.py tests/test_llm_memory_extractor.py tests/test_environment_extractor.py tests/test_tool_extractor.py tests/test_workflow_extractor.py tests/test_ingestion_to_extractors.py -q +``` + +这些测试使用 fake LLM client,不需要真实 API key,也不需要下载模型。 + +## 9. 对外汇报话术 + +可以这样讲: + +> 我这次把 B 侧抽取方向从规则增强改成了 LLM-only。旧的函数名和数据结构没有变, +> 但内部不再依赖正则、关键词、频次统计或长度阈值。现在输入统一是 MemoryEvent, +> 入模前先脱敏,LLM 返回结构化 JSON,最后由本地 validator/merger 转成 +> MemoryCandidate。这样既满足主管说的“大模型替代硬编码”,又不破坏团队已有接口。 diff --git a/docs/B_MEMORY_ADMISSION_CN.md b/docs/B_MEMORY_ADMISSION_CN.md new file mode 100644 index 0000000..be5b971 --- /dev/null +++ b/docs/B_MEMORY_ADMISSION_CN.md @@ -0,0 +1,117 @@ +# B 侧记忆准入与正式入库设计 + +## 1. 任务边界 + +本模块承接 B 侧抽取器输出的 `MemoryCandidate`,将经过验证和冲突判断的候选转成正式 `MemoryRecord`。 + +完整链路为: + +```text +MemoryEvent + -> LLM-only extractors + -> MemoryCandidate + -> MemoryAdmissionService + -> MemoryRecord + -> SQLite / 后续向量库 +``` + +本实现不修改已有抽取器函数签名、不修改核心字段名、不删除枚举,也不改变 Phase 0 SQLite 表结构。 + +## 2. LLM 与程序的职责划分 + +LLM 负责语义判断: + +- `create`:没有等价或冲突的已有记忆,创建正式记忆。 +- `duplicate`:与已有记忆语义相同,不重复写入。 +- `merge`:内容互补,输出合并后的标准内容并创建新版本。 +- `replace`:新事实更新或否定旧事实,新版本生效、旧版本被替代。 +- `coexist`:场景或作用域不同,两条记忆同时有效。 +- `pending`:证据不足或需要人工确认。 +- `reject`:不应进入长期记忆。 + +确定性程序负责工程安全: + +- 校验 LLM JSON 结构和目标记忆 ID。 +- 入模前遮盖密钥等敏感字段。 +- 拒绝敏感或不完整候选。 +- 生成稳定 `memory_id`,保证重试幂等。 +- 管理版本号和合法生命周期状态流转。 +- 使用 SQLite 写事务保证冲突处理和入库的原子性。 +- 并发期间发现快照变化时重新执行语义判断。 +- LLM 不可用时保持 `PENDING`,不进行猜测性写入。 + +代码中不存在以文本长度、关键词或固定置信度阈值替代语义判断的入库规则。 + +## 3. 生命周期 + +准入层支持以下主要状态流转: + +```text +PENDING -> ACTIVE / REJECTED / DELETED +ACTIVE -> SUPERSEDED / ARCHIVED / EXPIRED / DELETED +ARCHIVED -> ACTIVE / DELETED +SUPERSEDED -> ARCHIVED / DELETED +EXPIRED -> ARCHIVED / DELETED +REJECTED -> DELETED +``` + +相同状态的重复写入视为幂等操作;其他非法流转会抛出 `InvalidMemoryTransition`。 + +## 4. 冲突与版本处理 + +冲突判断输入包括: + +- 当前 `MemoryCandidate` 完整结构。 +- 同一 `user_id + memory_type + key` 下所有 `ACTIVE` 记忆。 +- 不同 `scenario` 的记录也会提供给 LLM,用于判断是否可以并存。 + +`merge` 和 `replace` 会: + +1. 校验模型引用的目标仍然处于 `ACTIVE`。 +2. 将目标记录更新为 `SUPERSEDED`。 +3. 创建版本号递增的新 `ACTIVE` 记录。 +4. 在本次返回的 `MemoryRecord.supersedes` 中给出被替代记录 ID。 + +Phase 0 表没有 `supersedes`、`metadata`、`tags` 等列,因此这些扩展信息目前只保留在返回对象中;持久层仍严格使用原表字段。后续若团队统一允许扩表,再补充完整的持久化追溯关系。 + +## 5. 并发和可用性 + +- SQLite 使用 WAL 模式,读写可并行。 +- 正式写入使用 `BEGIN IMMEDIATE`,冲突更新和新版本插入在同一事务完成。 +- 调用 LLM 时不持有数据库写锁,避免模型延迟阻塞其他请求。 +- 写入前重新核对活跃记忆快照;快照变化会重新读取并重新让 LLM 判断。 +- `memory_id` 根据候选身份、场景和最终内容稳定生成,相同事实重试不会生成重复行。 +- LLM 调用异常时返回 `PENDING` 且不修改数据库,防止错误覆盖已有记忆。 + +## 6. 文件说明 + +- `memory/admission.py`:准入服务、SQLite 事务和生命周期操作入口。 +- `memory/conflict_resolver.py`:LLM 冲突 prompt、JSON Schema 和输出校验。 +- `memory/version_manager.py`:稳定记忆 ID 和版本号计算。 +- `memory/lifecycle_state.py`:合法状态流转定义。 +- `tests/test_memory_admission.py`:正式入库、冲突、异常和并发测试。 +- `tests/test_conflict_resolver.py`:LLM 决策契约测试。 +- `demo/b_memory_admission_demo.py`:投屏展示的入库前后数据。 + +## 7. 本地运行 + +只看入库前后效果: + +```powershell +cd E:\OS +python demo\b_memory_admission_demo.py +``` + +运行新增测试: + +```powershell +python -m pytest tests\test_conflict_resolver.py tests\test_memory_admission.py -v +``` + +运行入库相关回归测试: + +```powershell +python -m pytest tests\test_conflict_resolver.py tests\test_memory_admission.py tests\test_memory_store.py -v +``` + +演示和测试均使用 Fake LLM,不访问外部模型,也不需要额外安装依赖。生产环境只需要注入实现 `complete_json(prompt, schema)` 的真实 LLM 客户端。 diff --git a/docs/B_MEMORY_ADMISSION_LOCAL_TEST_REPORT_CN.md b/docs/B_MEMORY_ADMISSION_LOCAL_TEST_REPORT_CN.md new file mode 100644 index 0000000..86435cb --- /dev/null +++ b/docs/B_MEMORY_ADMISSION_LOCAL_TEST_REPORT_CN.md @@ -0,0 +1,118 @@ +# B 侧记忆准入功能本地测试报告 + +## 基本信息 + +- 测试日期:2026-08-02 +- 本地分支:`feat/b-memory-admission` +- 基线:`feat/b-llm-only-memory-extractors`(commit `55438cf`) +- 提交流程:本地验证通过后推送个人 fork,并向团队 `master` 提交独立 PR +- Python:3.13.5 +- pytest:8.3.4 + +## 本次实现 + +新增 `MemoryCandidate -> MemoryRecord` 正式入库链路,覆盖: + +- LLM 语义准入和冲突判断。 +- 创建、语义去重、合并、替换、场景并存、待确认和拒绝。 +- 稳定 ID 与重复请求幂等。 +- 版本递增和旧版本 `SUPERSEDED`。 +- 合法生命周期状态流转。 +- SQLite WAL、写事务、并发快照校验和事务回滚。 +- 敏感信息保护、模型输出结构校验和模型异常时的失败关闭。 +- `MemoryEvent -> LLM extractor -> MemoryCandidate -> MemoryRecord` 对接测试。 + +未修改核心模型字段、现有函数签名、已有枚举值和 Phase 0 表结构。 + +## 测试结果 + +### 1. 入库专项与原存储回归 + +```powershell +python -m pytest tests/test_conflict_resolver.py tests/test_memory_admission.py tests/test_memory_store.py -q +``` + +结果:`62 passed`。 + +### 2. B 侧相关链路回归 + +```powershell +python -m pytest tests/test_conflict_resolver.py tests/test_memory_admission.py tests/test_memory_store.py tests/test_extractors.py tests/test_llm_memory_extractor.py tests/test_ingestion_to_extractors.py -q +``` + +结果:`81 passed`。 + +### 3. 新增入库模块覆盖率 + +```powershell +python -m coverage run -m pytest tests/test_conflict_resolver.py tests/test_memory_admission.py tests/test_memory_store.py -q +python -m coverage report -m --include="memory/admission.py,memory/conflict_resolver.py,memory/version_manager.py,memory/lifecycle_state.py" +``` + +结果: + +| 模块 | 覆盖率 | +| --- | ---: | +| `memory/admission.py` | 89% | +| `memory/conflict_resolver.py` | 90% | +| `memory/lifecycle_state.py` | 100% | +| `memory/version_manager.py` | 100% | +| 合计 | 90% | + +满足原任务要求的覆盖率不低于 85%。 + +### 4. 全量测试 + +```powershell +python -m pytest tests -q --tb=short --disable-warnings +``` + +结果:`179 passed, 5 failed`。 + +5 项失败全部来自 `tests/test_flow.py`,共同原因是仓库缺少: + +```text +E:\OS\data\raw\office_demo_events.jsonl +``` + +该问题在本次开发前已经存在,与新增入库代码无关。本次未擅自创建或修改 A 侧演示数据。 + +## 关键测试场景 + +- 首次候选正常生成 `ACTIVE/version=1` 正式记忆。 +- 同一候选重复提交不重复入库,也不重复调用 LLM。 +- 语义相同但表述不同的候选复用已有记录。 +- 新偏好替换旧偏好:旧记录 `SUPERSEDED`,新记录 `ACTIVE/version=2`。 +- 信息互补时生成合并后的新版本。 +- 不同场景的记忆允许同时保持 `ACTIVE`。 +- LLM 输出非法或要求人工审核时进入 `PENDING`。 +- LLM 不可用时不修改数据库。 +- 敏感候选在调用 LLM 前被拒绝。 +- 两个线程同时提交同一候选时只产生一条记录。 +- 替换过程中模拟写入失败时,事务回滚且旧记忆仍保持 `ACTIVE`。 +- 生命周期非法状态流转会被拒绝。 + +## 演示效果 + +```powershell +python demo\b_memory_admission_demo.py +``` + +演示会依次打印: + +1. 入库前的 `MemoryCandidate`。 +2. 脱敏后发送给 LLM 的候选和已有记忆。 +3. Fake LLM 的原始结构化决策。 +4. 入库后的 `AdmissionResult` 和正式 `MemoryRecord`。 +5. PDF 偏好被 Word 偏好替换后的两个版本。 +6. 相同候选重试的幂等结果。 +7. `ACTIVE -> ARCHIVED -> ACTIVE` 生命周期流转。 + +演示使用临时 SQLite 和 Fake LLM,不访问外部服务,退出后自动删除临时数据库。 + +## 当前限制与对接点 + +- Phase 0 `memories` 表没有 `source_events`、`tags`、`metadata`、`supersedes` 和 `vector_id` 列,因此这些扩展信息只在本次返回对象中完整保留,重新从 SQLite 读取时只能恢复原表字段。 +- 当前正式写入目标为 SQLite;仓库现有向量存储实现仍为空,待 C 侧提供稳定适配器后再接入。 +- D 侧草稿 PR #4 修改的是入库后的生命周期、整合和遗忘流程;本分支使用新的 `memory/admission.py` 与 `memory/lifecycle_state.py`,避免直接覆盖 D 侧 `memory/lifecycle.py`。 +- 合并前应以届时最新 `master` 为基线重新执行全量测试和 B-D 对接测试。 diff --git a/extractors/__init__.py b/extractors/__init__.py index 2fa8037..99df84e 100644 --- a/extractors/__init__.py +++ b/extractors/__init__.py @@ -1,7 +1,20 @@ """Extraction strategies for memory candidates.""" from .environment_extractor import EnvironmentExtractor +from .llm_memory_extractor import CandidateMerger, CandidateValidator, HybridMemoryExtractor, LLMMemoryExtractor from .knowledge_extractor import KnowledgeExtractor from .preference_extractor import PreferenceExtractor from .tool_extractor import ToolExtractor from .workflow_extractor import WorkflowExtractor + +__all__ = [ + "CandidateMerger", + "CandidateValidator", + "EnvironmentExtractor", + "HybridMemoryExtractor", + "KnowledgeExtractor", + "LLMMemoryExtractor", + "PreferenceExtractor", + "ToolExtractor", + "WorkflowExtractor", +] diff --git a/extractors/common.py b/extractors/common.py new file mode 100644 index 0000000..5f6f9ff --- /dev/null +++ b/extractors/common.py @@ -0,0 +1,46 @@ +"""Shared infrastructure helpers for extractor implementations. + +This file intentionally contains no memory extraction rules. It only provides +deterministic identifiers, key normalization, and logger creation. +""" + +from __future__ import annotations + +import hashlib +import logging +import re +from typing import Any + +from core.constants import MemoryType + + +DEFAULT_ID_LENGTH = 32 +SHORT_ID_LENGTH = 16 +HASH_TOKEN_LENGTH = 12 + +_SLUG_RE = re.compile(r"[^0-9a-zA-Z\u4e00-\u9fff]+") + + +def slugify(value: Any, *, fallback: str = "memory") -> str: + token = _SLUG_RE.sub("_", str(value or "").strip().lower()).strip("_") + token = re.sub(r"_+", "_", token) + return token or fallback + + +def stable_candidate_id( + user_id: str, + memory_type: MemoryType, + key: str, + *, + length: int = DEFAULT_ID_LENGTH, +) -> str: + payload = f"{user_id}\x1f{memory_type.value}\x1f{key}".encode("utf-8") + return hashlib.sha256(payload).hexdigest()[:length] + + +def stable_hash_token(value: str, *, length: int = HASH_TOKEN_LENGTH) -> str: + return hashlib.sha256(value.encode("utf-8")).hexdigest()[:length] + + +def extractor_logger(module_name: str) -> logging.Logger: + return logging.getLogger(module_name) diff --git a/extractors/environment_extractor.py b/extractors/environment_extractor.py index c60c6b6..1dea2c1 100644 --- a/extractors/environment_extractor.py +++ b/extractors/environment_extractor.py @@ -1,309 +1,43 @@ -from __future__ import annotations - -import hashlib -import re -from collections.abc import Iterator -from typing import Any - -from core.constants import MemoryType, Scene -from core.models import MemoryCandidate - - -_SENSITIVE_TOKENS = ("password", "passwd", "secret", "token", "api_key", "access_key", "credential", "authorization") -_DOWNLOAD_KEYS = {"downloads", "download", "download_dir", "downloads_dir", "download_path"} -_DOCUMENT_KEYS = {"documents", "document", "docs", "doc", "documents_dir", "document_dir", "docs_dir", "doc_dir", "documents_path", "docs_path"} -_LANGUAGE_KEYS = {"language", "lang", "system_language", "ui_language", "display_language"} -_REGION_KEYS = {"region", "country", "locale_region", "system_region"} -_LOCALE_KEYS = {"locale", "system_locale", "language_locale"} -_SOFTWARE_KEYS = {"installed_software", "installed_applications", "applications", "software", "packages", "installed_packages"} -_VERSION_KEYS = {"os_version", "system_version", "operating_system", "platform", "os", "system"} - - -def _slugify(value: str) -> str: - token = re.sub(r"[^0-9a-zA-Z\u4e00-\u9fff]+", "_", value.strip().lower()) - return token.strip("_") or "environment" - - -def _stable_candidate_id(user_id: str, key: str) -> str: - payload = f"{user_id}\x1f{MemoryType.ENVIRONMENT.value}\x1f{key}".encode("utf-8") - return hashlib.sha256(payload).hexdigest()[:32] - - -def _walk(payload: Any, path: tuple[str, ...] = (), visited: set[int] | None = None) -> Iterator[tuple[tuple[str, ...], Any]]: - """Yield nested values without recursing forever on malformed cyclic input.""" - visited = visited if visited is not None else set() - if isinstance(payload, dict): - marker = id(payload) - if marker in visited: - return - visited.add(marker) - for raw_key, value in payload.items(): - key = str(raw_key).strip().lower() - if any(token in key for token in _SENSITIVE_TOKENS): - continue - yield from _walk(value, path + (key,), visited) - visited.remove(marker) - return - if isinstance(payload, (list, tuple, set)): - marker = id(payload) - if marker in visited: - return - visited.add(marker) - values = sorted(payload, key=str) if isinstance(payload, set) else payload - for index, value in enumerate(values): - yield from _walk(value, path + (str(index),), visited) - visited.remove(marker) - return - yield path, payload - - -def _leaf_key(path: tuple[str, ...]) -> str: - return path[-1] if path else "" - - -def _sanitise_path(value: str, directory: str) -> str: - raw_path = value.strip() - path = raw_path.replace("/", "\\") - windows_home = re.compile(r"^[a-zA-Z]:\\users\\[^\\]+(?P\\.*)$", re.IGNORECASE) - unix_home = re.compile(r"^/(?:home|users)/[^/]+(?P/.*)$", re.IGNORECASE) - match = windows_home.match(path) - if match: - path = "~" + match.group("tail") - else: - unix_match = unix_home.match(raw_path) - if unix_match: - path = "~" + unix_match.group("tail").replace("/", "\\") - path_tail = path.rstrip("\\").rsplit("\\", 1)[-1].lower() - expected_tails = { - "downloads": {"downloads", "download"}, - "documents": {"documents", "document", "docs", "doc"}, - }.get(directory.lower(), {directory.lower()}) - if path_tail not in expected_tails: - if re.match(r"^(?:[A-Za-z]:\\|\\\\|\\|~\\)", path) or raw_path.startswith(("/", "~")): - return path - return directory - return path +"""Environment extraction facade backed entirely by the configured LLM. +Environment memory is no longer inferred with path/category keyword rules. +The tool output is wrapped as a ``SYSTEM_CONTEXT`` event and sent to the +configured LLM client. +""" -def _path_category(path: tuple[str, ...], value: Any) -> str | None: - key = _leaf_key(path) - if key in _DOWNLOAD_KEYS: - return "downloads" - if key in _DOCUMENT_KEYS: - return "documents" - if isinstance(value, str): - lowered = value.replace("/", "\\").rstrip("\\").lower() - if lowered.endswith("\\downloads"): - return "downloads" - if lowered.endswith("\\documents") or lowered.endswith("\\docs"): - return "documents" - return None - +from __future__ import annotations -def _parse_locale(value: str) -> tuple[str | None, str | None]: - normalized = value.strip().replace("-", "_").split(".", 1)[0] - parts = [part for part in normalized.split("_") if part] - if not parts: - return None, None - language = parts[0].lower() - region = parts[1].upper() if len(parts) > 1 and len(parts[1]) in {2, 3} else None - return language, region +from datetime import datetime, timezone +from core.constants import EventType, MemoryType, Scene +from core.models import MemoryCandidate, MemoryEvent -def _software_entries(value: Any) -> list[tuple[str, str | None]]: - entries: list[tuple[str, str | None]] = [] - if isinstance(value, str): - for item in re.split(r"[,;\n]", value): - name = item.strip() - if name: - entries.append((name, None)) - elif isinstance(value, dict): - if any(key in value for key in ("name", "display_name", "package")): - name = value.get("name") or value.get("display_name") or value.get("package") - version = value.get("version") or value.get("release") - if isinstance(name, str) and name.strip(): - entries.append((name.strip(), str(version).strip() if version is not None else None)) - else: - for name, version in value.items(): - if isinstance(version, dict): - nested_version = version.get("version") or version.get("release") - entries.append((str(name), str(nested_version).strip() if nested_version is not None else None)) - elif version is not None: - entries.append((str(name), str(version).strip())) - elif isinstance(value, (list, tuple, set)): - values = sorted(value, key=str) if isinstance(value, set) else value - for item in values: - entries.extend(_software_entries(item)) - return [(name, version) for name, version in entries if name.strip()] +from .llm_memory_extractor import extract_candidates_with_default_client -def _make_candidate( - *, - user_id: str, - key: str, - content: str, - confidence: float, - metadata: dict[str, Any], -) -> MemoryCandidate: - return MemoryCandidate( - candidate_id=_stable_candidate_id(user_id, key), - user_id=user_id, - memory_type=MemoryType.ENVIRONMENT, - key=key, - content=content, +def _event_from_tool_output(output: dict) -> MemoryEvent: + return MemoryEvent( + event_id=str(output.get("event_id") or "synthetic-environment-output"), + raw_event_id=str(output.get("raw_event_id") or output.get("event_id") or "synthetic-environment-output"), + user_id=str(output.get("user_id") or ""), + session_id=str(output.get("session_id") or ""), + task_id=str(output.get("task_id") or ""), + event_type=EventType.SYSTEM_CONTEXT, scenario=Scene.SYSTEM, - confidence=confidence, - source="environment_tool_output", - tags=["environment", metadata["category"]], - metadata=metadata, + source="tool_output", + actor="system", + output=output, + timestamp=datetime.now(timezone.utc), ) class EnvironmentExtractor: @staticmethod def extract_from_tool_output(output: dict) -> list[MemoryCandidate]: - """Extract normalised, non-sensitive environment facts from tool output.""" - if not isinstance(output, dict): - raise TypeError("output must be a dict") - - user_id = str(output.get("user_id", "")).strip() - candidates: dict[str, MemoryCandidate] = {} - - def add(candidate: MemoryCandidate) -> None: - existing = candidates.get(candidate.key) - if existing is None or candidate.confidence > existing.confidence: - candidates[candidate.key] = candidate - - for path, value in _walk(output): - key = _leaf_key(path) - if value is None: - continue - - directory = _path_category(path, value) - if directory and isinstance(value, str) and value.strip(): - normalised_path = _sanitise_path(value, directory) - add( - _make_candidate( - user_id=user_id, - key=f"environment.path.{directory}", - content=f"常用目录 {directory}: {normalised_path}", - confidence=0.95, - metadata={"category": "path", "directory": directory, "path": normalised_path, "source_key": ".".join(path)}, - ) - ) - - if key in (_LANGUAGE_KEYS | _LOCALE_KEYS) and isinstance(value, str) and value.strip(): - language, inferred_region = _parse_locale(value) - if language: - add( - _make_candidate( - user_id=user_id, - key="environment.locale.language", - content=f"系统语言: {language}", - confidence=0.9, - metadata={"category": "locale", "language": language, "source_key": ".".join(path)}, - ) - ) - if inferred_region: - add( - _make_candidate( - user_id=user_id, - key="environment.locale.region", - content=f"系统地区: {inferred_region}", - confidence=0.82, - metadata={"category": "locale", "region": inferred_region, "source_key": ".".join(path)}, - ) - ) - - if key in _REGION_KEYS and isinstance(value, str) and value.strip(): - region = value.strip().upper() - add( - _make_candidate( - user_id=user_id, - key="environment.locale.region", - content=f"系统地区: {region}", - confidence=0.9, - metadata={"category": "locale", "region": region, "source_key": ".".join(path)}, - ) - ) - - # Software collections need to retain their container shape, so inspect - # top-level and nested dictionaries separately from scalar walking. - containers: list[tuple[tuple[str, ...], Any]] = [] - - def collect_containers(payload: Any, path: tuple[str, ...] = (), visited: set[int] | None = None) -> None: - visited = visited if visited is not None else set() - if not isinstance(payload, dict): - return - marker = id(payload) - if marker in visited: - return - visited.add(marker) - for raw_key, value in payload.items(): - key = str(raw_key).strip().lower() - child_path = path + (key,) - if any(token in key for token in _SENSITIVE_TOKENS): - continue - if key in _SOFTWARE_KEYS: - containers.append((child_path, value)) - if isinstance(value, dict): - collect_containers(value, child_path, visited) - visited.remove(marker) - - collect_containers(output) - for path, value in containers: - for name, version in _software_entries(value): - software_key = f"environment.software.{_slugify(name)}" - display = f"已安装软件: {name}" + (f" {version}" if version else "") - add( - _make_candidate( - user_id=user_id, - key=software_key, - content=display, - confidence=0.9, - metadata={"category": "software", "name": name, "version": version, "source_key": ".".join(path)}, - ) - ) - - version_values: list[tuple[tuple[str, ...], str]] = [] - - def collect_versions(payload: Any, path: tuple[str, ...] = (), visited: set[int] | None = None) -> None: - visited = visited if visited is not None else set() - if not isinstance(payload, dict): - return - marker = id(payload) - if marker in visited: - return - visited.add(marker) - for raw_key, value in payload.items(): - key = str(raw_key).strip().lower() - child_path = path + (key,) - if any(token in key for token in _SENSITIVE_TOKENS): - continue - if key in _VERSION_KEYS: - if isinstance(value, str) and value.strip(): - version_values.append((child_path, value.strip())) - elif isinstance(value, dict): - name = value.get("name") or value.get("os_name") or value.get("distribution") - version = value.get("version") or value.get("release") or value.get("build") - rendered = " ".join(str(item).strip() for item in (name, version) if item is not None and str(item).strip()) - if rendered: - version_values.append((child_path, rendered)) - if isinstance(value, dict): - collect_versions(value, child_path, visited) - visited.remove(marker) - - collect_versions(output) - if version_values: - path, version = sorted(version_values, key=lambda item: (len(item[0]), item[0]))[0] - add( - _make_candidate( - user_id=user_id, - key="environment.system.version", - content=f"系统版本: {version}", - confidence=0.9, - metadata={"category": "system", "version": version, "source_key": ".".join(path)}, - ) - ) - - return [candidates[key] for key in sorted(candidates)] + if not isinstance(output, dict) or not output: + return [] + return extract_candidates_with_default_client( + [_event_from_tool_output(output)], + mode="environment_from_tool_output", + memory_types={MemoryType.ENVIRONMENT}, + ) diff --git a/extractors/knowledge_extractor.py b/extractors/knowledge_extractor.py index ff24473..820ec43 100644 --- a/extractors/knowledge_extractor.py +++ b/extractors/knowledge_extractor.py @@ -1,613 +1,40 @@ -from __future__ import annotations - -import hashlib -import re -from collections import defaultdict -from dataclasses import dataclass -from typing import Any, Iterable - -from core.constants import MemoryType, Scene -from core.models import MemoryCandidate, MemoryEvent - -_SENSITIVE_TOKENS = ( - "password", - "passwd", - "secret", - "token", - "api_key", - "apikey", - "access_key", - "credential", - "authorization", -) -_SECRET_ASSIGNMENT_RE = re.compile( - r"(?i)\b(api[_-]?key|password|passwd|secret|token|access[_-]?key|authorization|credential)\b\s*[:=]\s*['\"]?[^'\"\s,;]+" -) -_BEARER_RE = re.compile(r"(?i)\bbearer\s+[A-Za-z0-9._~+/=-]+") -_LONG_SECRET_RE = re.compile(r"(?.+?))(?:(?:\r?\n)|[。;;])\s*(?:(?:答案|A|解决方案|解决办法|处理方式)[::]\s*(?P.+))", - re.IGNORECASE | re.DOTALL, - ), - re.compile( - r"(?:(?P[^。!??]*?(?:怎么|如何|为何|为什么|怎样|报错|失败|无法|不能)[^。!??]*?))(?:(?:。|!|\?|?))?\s*(?:(?P[^。!??]*?(?:先|然后|再|最后|可以|建议|检查|重试|修复|解决)[^。!??]*))", - re.IGNORECASE | re.DOTALL, - ), -] - -_GUIDE_HINTS = ( - "步骤", - "教程", - "指南", - "配置", - "设置", - "安装", - "部署", - "初始化", - "运行", - "导出", - "合并", -) - -_TEMPLATE_DOMAIN_RULES = [ - ( - "batch_export", - ( - "batch export", - "batch_export", - "批量导出", - "导出", - "export", - "save as", - ), - ), - ( - "merge_files", - ( - "merge files", - "merge", - "combine", - "合并文件", - "合并", - "拼接", - ), - ), - ( - "desktop_config", - ( - "desktop config", - "desktop", - "桌面", - "配置", - "settings", - "setup desktop", - ), - ), - ( - "software_setup", - ( - "software setup", - "setup", - "install", - "安装", - "部署", - "初始化", - "environment setup", - ), - ), -] - - -def _slugify(text: str) -> str: - token = re.sub(r"[^0-9a-zA-Z\u4e00-\u9fff]+", "_", text.strip().lower()) - token = re.sub(r"_+", "_", token).strip("_") - return token or "knowledge" - - -def _is_sensitive_key(key: Any) -> bool: - normalized = str(key).strip().lower().replace("-", "_") - return any(token in normalized for token in _SENSITIVE_TOKENS) - - -def _redact_sensitive_text(text: str) -> str: - text = _SECRET_ASSIGNMENT_RE.sub(lambda match: f"{match.group(1)}=", text) - text = _BEARER_RE.sub("Bearer ", text) - return _LONG_SECRET_RE.sub("", text) - - -def _stable_candidate_id(user_id: str, memory_type: MemoryType, key: str) -> str: - """Return a deterministic candidate id for idempotent extraction output.""" - payload = f"{user_id}\x1f{memory_type.value}\x1f{key}".encode("utf-8") - return hashlib.sha256(payload).hexdigest()[:32] - - -def _text_fragments(payload: Any) -> list[str]: - if payload is None: - return [] - if isinstance(payload, str): - text = _redact_sensitive_text(payload.strip()) - return [text] if text else [] - if isinstance(payload, (int, float, bool)): - return [str(payload)] - if isinstance(payload, dict): - fragments: list[str] = [] - for key, value in payload.items(): - if key in {"event_id", "raw_event_id", "timestamp", "created_at", "updated_at"} or _is_sensitive_key(key): - continue - fragments.extend(_text_fragments(value)) - return fragments - if isinstance(payload, (list, tuple, set)): - fragments: list[str] = [] - items = sorted(payload, key=lambda item: str(item)) if isinstance(payload, set) else payload - for item in items: - fragments.extend(_text_fragments(item)) - return fragments - return [str(payload)] - - -def _flatten_event_text(event: MemoryEvent) -> str: - parts = [ - event.content or "", - event.tool_name or "", - *(_text_fragments(event.input)), - *(_text_fragments(event.output)), - *(_text_fragments(event.metadata)), - ] - return _redact_sensitive_text("\n".join(part for part in parts if part and str(part).strip())) - - -def _normalize_for_template(text: str) -> str: - value = text.lower() - value = re.sub(r"https?://\S+", "", value) - value = re.sub(r"[a-zA-Z]:\\[^\s]+", "", value) - value = re.sub(r"/[^\s]+", "", value) - value = re.sub( - r"\b[\w.-]*\d[\w.-]*\.(?:csv|xlsx|xls|txt|json|jsonl|docx|pdf|pptx|md|zip)\b", - "", - value, - ) - value = re.sub(r"\b[\w.-]*\d[\w.-]*\b", "", value) - value = re.sub(r"\b[0-9]+\b", "", value) - value = re.sub(r"\b[0-9a-f]{8,}\b", "", value) - value = re.sub(r"['\"][^'\"]+['\"]", "", value) - value = re.sub(r"\s+", " ", value) - return value.strip() - - -def _summarize_mapping(payload: dict[str, Any] | None, *, max_items: int = 4) -> str: - if not payload: - return "" - fragments: list[str] = [] - for key, value in payload.items(): - if _is_sensitive_key(key): - continue - if value is None: - continue - if isinstance(value, dict): - continue - if isinstance(value, (list, tuple, set)): - normalized = ", ".join(str(item) for item in value if item is not None) - else: - normalized = _redact_sensitive_text(str(value)) - if normalized: - fragments.append(f"{key}={normalized}") - if len(fragments) >= max_items: - break - return "; ".join(fragments) - - -def _make_candidate( - *, - user_id: str, - memory_type: MemoryType, - key: str, - content: str, - scenario: Scene | str, - confidence: float, - source: str, - source_events: list[str], - source_summaries: list[str], - tags: list[str], - metadata: dict[str, Any] | None = None, -) -> MemoryCandidate: - safe_content = _redact_sensitive_text(content) - safe_summaries = [_redact_sensitive_text(summary) for summary in source_summaries] - return MemoryCandidate( - candidate_id=_stable_candidate_id(user_id, memory_type, key), - user_id=user_id, - memory_type=memory_type, - key=key, - content=safe_content, - scenario=scenario, - confidence=max(0.5, min(confidence, 1.0)), - source=source, - source_events=source_events, - source_summaries=safe_summaries, - tags=tags, - metadata=metadata or {}, - ) - - -def _ensure_event_context(candidate: MemoryCandidate, event: MemoryEvent, *, source: str) -> MemoryCandidate: - candidate.user_id = event.user_id - candidate.scenario = event.scenario - candidate.source = source - if event.event_id not in candidate.source_events: - candidate.source_events.append(event.event_id) - if not candidate.source_summaries and event.content: - candidate.source_summaries = [event.content] - return candidate - - -def _extract_tool_intent(event: MemoryEvent) -> str: - fragments = " ".join( - part.lower() - for part in ( - event.tool_name or "", - event.content or "", - _summarize_mapping(event.input), - _summarize_mapping(event.output), - ) - if part - ) - - for intent, keywords in _TEMPLATE_DOMAIN_RULES: - if any(keyword in fragments for keyword in keywords): - return intent - return "tool_case" - +"""Knowledge/template extraction facade backed entirely by the configured LLM. -def _completeness_score(*, has_input: bool, has_output: bool, has_content: bool, has_metadata: bool) -> float: - score = 0.5 - if has_input: - score += 0.15 - if has_output: - score += 0.2 - if has_content: - score += 0.1 - if has_metadata: - score += 0.05 - return min(score, 0.95) +The public API is unchanged, but extraction no longer uses FAQ regexes, +keyword triggers, or repeated-task heuristics. The configured LLM is +responsible for deciding whether a tool result, conversation, or event batch +contains reusable knowledge or templates. +""" +from __future__ import annotations -def _extract_faq_candidates(event: MemoryEvent, text: str) -> list[MemoryCandidate]: - candidates: list[MemoryCandidate] = [] - seen: set[str] = set() - - for pattern in _FAQ_PATTERNS: - for match in pattern.finditer(text): - question = (match.groupdict().get("question") or "").strip() - answer = (match.groupdict().get("answer") or "").strip() - if not question or not answer: - continue - signature = _normalize_for_template(f"{question} :: {answer}") - if signature in seen: - continue - seen.add(signature) - topic = _slugify(question[:24]) - confidence = 0.88 - if len(question) > 18: - confidence += 0.04 - if len(answer) > 18: - confidence += 0.04 - candidates.append( - _make_candidate( - user_id=event.user_id, - memory_type=MemoryType.KNOWLEDGE, - key=f"knowledge.faq.{topic}", - content=f"问题:{question};解决:{answer}", - scenario=event.scenario, - confidence=confidence, - source="knowledge_conversation", - source_events=[event.event_id], - source_summaries=[question, answer], - tags=["faq", "solution"], - metadata={ - "topic": topic, - "question": question, - "answer": answer, - "kind": "faq_solution", - }, - ) - ) - - return candidates - - -def _extract_guide_candidates(event: MemoryEvent, text: str) -> list[MemoryCandidate]: - lowered = text.lower() - guide_triggers = ("怎么", "如何", "步骤", "配置", "安装", "设置", "导出", "合并", "教程", "指南") - if not any(trigger in lowered for trigger in guide_triggers) and not any(hint in text for hint in _GUIDE_HINTS): - return [] - - steps = [ - line.strip(" -•\t") - for line in re.split(r"[\n\r;;。]+", text) - if line.strip() - ] - step_like = [step for step in steps if re.match(r"^(?:先|然后|再|最后|步骤|\d+[).、-])", step)] - content = ";".join(step_like[:5] if step_like else steps[:3]) - if not content: - content = text.strip() - topic_seed = step_like[0] if step_like else steps[0] - topic = _slugify(topic_seed[:24]) - confidence = 0.72 - if step_like: - confidence += min(0.15, 0.03 * len(step_like)) - if len(steps) >= 3: - confidence += 0.05 - - return [ - _make_candidate( - user_id=event.user_id, - memory_type=MemoryType.KNOWLEDGE, - key=f"knowledge.guide.{topic}", - content=f"操作指南:{content}", - scenario=event.scenario, - confidence=confidence, - source="knowledge_conversation", - source_events=[event.event_id], - source_summaries=[content], - tags=["guide", "howto"], - metadata={ - "topic": topic, - "step_count": len(step_like) if step_like else len(steps), - "kind": "operation_guide", - }, - ) - ] - - -def _extract_system_guide(event: MemoryEvent) -> list[MemoryCandidate]: - fragments = " ".join( - part.lower() - for part in ( - event.content or "", - _summarize_mapping(event.input), - _summarize_mapping(event.output), - _summarize_mapping(event.metadata), - ) - if part - ) - if not any(keyword in fragments for keyword in ("desktop", "桌面", "config", "配置", "settings", "安装", "setup", "software")): - return [] - - topic = "system_setup" - if "desktop" in fragments or "桌面" in fragments: - topic = "desktop_config" - elif "安装" in fragments or "setup" in fragments or "software" in fragments: - topic = "software_setup" - - content = event.content.strip() if event.content else "" - if not content: - content = "系统操作指南" - - confidence = 0.74 - if event.input: - confidence += 0.08 - if event.output: - confidence += 0.08 - if event.metadata: - confidence += 0.04 - - return [ - _make_candidate( - user_id=event.user_id, - memory_type=MemoryType.KNOWLEDGE, - key=f"knowledge.system.{topic}", - content=f"系统操作指南:{content}", - scenario=event.scenario, - confidence=confidence, - source="knowledge_conversation", - source_events=[event.event_id], - source_summaries=[content], - tags=["system", topic], - metadata={ - "topic": topic, - "kind": "system_guide", - }, - ) - ] - - -def _template_signature(event: MemoryEvent) -> tuple[str, str, str]: - text = _flatten_event_text(event) - normalized = _normalize_for_template(text) - domain = "generic" - for name, keywords in _TEMPLATE_DOMAIN_RULES: - if any(keyword in normalized for keyword in keywords): - domain = name - break - structure = [] - if event.tool_name: - structure.append(f"tool:{event.tool_name.strip().lower()}") - if event.input: - structure.append("input:" + ",".join(sorted(str(key).lower() for key in event.input.keys()))) - if event.output: - structure.append("output:" + ",".join(sorted(str(key).lower() for key in event.output.keys()))) - if event.metadata: - structure.append("meta:" + ",".join(sorted(str(key).lower() for key in event.metadata.keys()))) - structure_key = "|".join(structure) if structure else "plain" - text_key = re.sub(r"\s+", " ", normalized) - return domain, structure_key, text_key - - -def _template_content(domain: str, events: list[MemoryEvent]) -> str: - representative = events[0] - sample = representative.content or representative.tool_name or domain - if domain == "batch_export": - return f"批量导出模板:{sample}" - if domain == "merge_files": - return f"合并文件模板:{sample}" - if domain == "desktop_config": - return f"桌面配置模板:{sample}" - if domain == "software_setup": - return f"软件安装/配置模板:{sample}" - return f"通用模板:{sample}" - - -def _template_confidence(events: list[MemoryEvent]) -> float: - if not events: - return 0.5 - count = len(events) - rich_events = 0 - for event in events: - if event.input: - rich_events += 1 - if event.output: - rich_events += 1 - if event.content: - rich_events += 1 - if event.metadata: - rich_events += 1 - richness = min(1.0, rich_events / (count * 4)) - confidence = 0.55 - confidence += min(0.18, 0.06 * max(0, count - 1)) - confidence += 0.18 * richness - return min(confidence, 0.95) - +from core.constants import MemoryType +from core.models import MemoryCandidate, MemoryEvent -def _dedupe_candidates(candidates: list[MemoryCandidate]) -> list[MemoryCandidate]: - deduped: list[MemoryCandidate] = [] - seen: set[tuple[str, str, str]] = set() - for candidate in candidates: - signature = (candidate.user_id, candidate.memory_type.value, candidate.key) - if signature in seen: - continue - seen.add(signature) - deduped.append(candidate) - return deduped +from .llm_memory_extractor import extract_candidates_with_default_client class KnowledgeExtractor: @staticmethod def extract_from_tool_result(event: MemoryEvent) -> list[MemoryCandidate]: - text = _flatten_event_text(event) - intent = _extract_tool_intent(event) - input_summary = _summarize_mapping(event.input) - output_summary = _summarize_mapping(event.output) - metadata_summary = _summarize_mapping(event.metadata) - - has_input = bool(event.input) - has_output = bool(event.output) - has_content = bool(event.content and event.content.strip()) - has_metadata = bool(event.metadata) - confidence = _completeness_score( - has_input=has_input, - has_output=has_output, - has_content=has_content, - has_metadata=has_metadata, + return extract_candidates_with_default_client( + [event], + mode="knowledge_from_tool_result", + memory_types={MemoryType.KNOWLEDGE, MemoryType.TEMPLATE}, ) - tool_name = _slugify(event.tool_name or "tool") - key = f"knowledge.tool_case.{intent}.{tool_name}" - content_bits = [f"工具用例:{event.tool_name or 'tool'}"] - if input_summary: - content_bits.append(f"输入 {input_summary}") - if output_summary: - content_bits.append(f"输出 {output_summary}") - elif text.strip(): - content_bits.append(f"结果 {text.strip()[:120]}") - - candidates = [ - _make_candidate( - user_id=event.user_id, - memory_type=MemoryType.KNOWLEDGE, - key=key, - content=";".join(content_bits), - scenario=event.scenario, - confidence=confidence, - source="knowledge_tool_result", - source_events=[event.event_id], - source_summaries=[summary for summary in (input_summary, output_summary, metadata_summary, event.content or "") if summary], - tags=["tool_case", intent], - metadata={ - "intent": intent, - "tool_name": event.tool_name, - "input_summary": input_summary, - "output_summary": output_summary, - "completeness": confidence, - }, - ) - ] - - if event.success is False or any(term in text.lower() for term in ("error", "exception", "failed", "失败", "报错", "异常")): - detail = event.content or output_summary or text.strip() - if detail: - candidates.append( - _make_candidate( - user_id=event.user_id, - memory_type=MemoryType.KNOWLEDGE, - key=f"knowledge.issue.{intent}.{tool_name}", - content=f"问题诊断:{detail}", - scenario=event.scenario, - confidence=max(0.62, confidence - 0.08), - source="knowledge_tool_result", - source_events=[event.event_id], - source_summaries=[detail], - tags=["issue", intent], - metadata={ - "intent": intent, - "tool_name": event.tool_name, - "kind": "problem_diagnosis", - }, - ) - ) - - return _dedupe_candidates([_ensure_event_context(candidate, event, source="knowledge_tool_result") for candidate in candidates]) - @staticmethod def extract_from_conversation(event: MemoryEvent) -> list[MemoryCandidate]: - text = _flatten_event_text(event) - candidates: list[MemoryCandidate] = [] - candidates.extend(_extract_faq_candidates(event, text)) - candidates.extend(_extract_guide_candidates(event, text)) - candidates.extend(_extract_system_guide(event)) - return _dedupe_candidates(candidates) + return extract_candidates_with_default_client( + [event], + mode="knowledge_from_conversation", + memory_types={MemoryType.KNOWLEDGE, MemoryType.TEMPLATE}, + ) @staticmethod def extract_templates(events: list[MemoryEvent]) -> list[MemoryCandidate]: - # Templates are personal memories. Grouping across users would merge - # evidence from different tenants and leak a second user's activity. - groups: dict[tuple[str, str, str, str], list[MemoryEvent]] = defaultdict(list) - for event in events: - domain, structure_key, text_key = _template_signature(event) - groups[(event.user_id, domain, structure_key, text_key)].append(event) - - candidates: list[MemoryCandidate] = [] - for (user_id, domain, structure_key, text_key), group in sorted(groups.items()): - if len(group) < 2: - continue - ordered_group = sorted(group, key=lambda event: (event.timestamp, event.event_id)) - representative = ordered_group[0] - confidence = _template_confidence(ordered_group) - template_key = ( - f"template.{domain}.{_slugify(structure_key)}." - f"{hashlib.sha256(text_key.encode('utf-8')).hexdigest()[:12]}" - ) - candidates.append( - _make_candidate( - user_id=user_id, - memory_type=MemoryType.TEMPLATE, - key=template_key, - content=_template_content(domain, ordered_group), - scenario=representative.scenario if all(event.scenario == representative.scenario for event in ordered_group) else Scene.GLOBAL, - confidence=confidence, - source="knowledge_template", - source_events=[event.event_id for event in ordered_group], - source_summaries=[event.content or _flatten_event_text(event)[:160] for event in ordered_group if event.content or _flatten_event_text(event)], - tags=["template", domain], - metadata={ - "domain": domain, - "structure_key": structure_key, - "count": len(ordered_group), - "kind": "reusable_template", - "normalized_text": text_key, - }, - ) - ) - - return _dedupe_candidates(candidates) + return extract_candidates_with_default_client( + events, + mode="template_from_session", + memory_types={MemoryType.TEMPLATE}, + ) diff --git a/extractors/llm_memory_extractor.py b/extractors/llm_memory_extractor.py new file mode 100644 index 0000000..f2ac3aa --- /dev/null +++ b/extractors/llm_memory_extractor.py @@ -0,0 +1,616 @@ +"""LLM-only semantic memory extraction pipeline for B-side extractors. + +This module is the single semantic extraction path for B-side memory work. +Deterministic regex/keyword extractors are intentionally not used as a +baseline here: production code injects an LLM JSON client, and tests inject a +fake client. Thin wrappers in ``preference_extractor.py``, +``knowledge_extractor.py``, ``workflow_extractor.py``, ``tool_extractor.py``, +and ``environment_extractor.py`` preserve the original public method +signatures while delegating extraction to this module. + +The remaining local logic is not memory extraction logic. It is limited to: + +- prompt packaging from ``MemoryEvent``; +- credential/contact redaction before prompt and after model output; +- schema conversion from model JSON to ``MemoryCandidate``; +- duplicate/conflict handling for already-produced candidates. +""" + +from __future__ import annotations + +import copy +import json +import re +from dataclasses import replace +from typing import Any, Iterable, Protocol + +from core.constants import MemoryType, Scene +from core.models import MemoryCandidate, MemoryEvent + +from .common import ( + SHORT_ID_LENGTH, + extractor_logger, + slugify as _shared_slugify, + stable_candidate_id as _shared_stable_candidate_id, +) + + +logger = extractor_logger(__name__) + + +class LLMJsonClient(Protocol): + """Adapter boundary for any JSON-returning LLM provider.""" + + def complete_json(self, prompt: str, schema: dict[str, Any]) -> Any: + """Return JSON-compatible data matching ``schema`` as closely as possible.""" + + +_DEFAULT_LLM_CLIENT: LLMJsonClient | None = None + + +def set_default_llm_client(client: LLMJsonClient | None) -> None: + """Set the process-local LLM client used by legacy extractor signatures.""" + + global _DEFAULT_LLM_CLIENT + _DEFAULT_LLM_CLIENT = client + + +def get_default_llm_client() -> LLMJsonClient | None: + """Return the process-local LLM client, if configured.""" + + return _DEFAULT_LLM_CLIENT + + +def clear_default_llm_client() -> None: + """Clear the process-local LLM client. + + Tests should call this in teardown to avoid cross-test leakage. + """ + + set_default_llm_client(None) + + +MEMORY_EXTRACTION_SCHEMA: dict[str, Any] = { + "type": "object", + "required": ["candidates"], + "properties": { + "candidates": { + "type": "array", + "items": { + "type": "object", + "required": [ + "is_memory_worthy", + "is_long_term", + "memory_type", + "category", + "value", + "content", + "confidence", + "evidence", + "reason", + ], + "properties": { + "is_memory_worthy": {"type": "boolean"}, + "is_long_term": {"type": "boolean"}, + "memory_type": { + "type": "string", + "enum": [memory_type.value for memory_type in MemoryType], + }, + "category": {"type": "string"}, + "value": {"type": "string"}, + "scope": {"type": "string"}, + "content": {"type": "string"}, + "confidence": {"type": "number", "minimum": 0, "maximum": 1}, + "evidence": {"type": "string"}, + "reason": {"type": "string"}, + "sensitivity": {"type": "string"}, + "conflicts_with": {"type": "array", "items": {"type": "string"}}, + }, + }, + } + }, +} + + +WORKFLOW_BOUNDARY_SCHEMA: dict[str, Any] = { + "type": "object", + "required": ["boundaries"], + "properties": { + "boundaries": { + "type": "array", + "items": { + "type": "object", + "required": ["start", "end", "reason"], + "properties": { + "start": {"type": "integer", "minimum": 0}, + "end": {"type": "integer", "minimum": 0}, + "reason": {"type": "string"}, + }, + }, + } + }, +} + + +_SECRET_KEY_RE = re.compile( + r"(api[_-]?key|password|passwd|secret|token|authorization|cookie|private[_-]?key)", + re.IGNORECASE, +) +_SECRET_VALUE_RE = re.compile( + r"(?i)\b(?:sk-[a-z0-9_-]+|Bearer\s+[a-z0-9._~+/-]+|" + r"[a-z0-9_-]{24,}\.[a-z0-9._-]+)\b" +) +_EMAIL_RE = re.compile(r"\b[\w.+-]+@[\w.-]+\.[a-zA-Z]{2,}\b") +_PHONE_RE = re.compile(r"(? str: + return _shared_stable_candidate_id(user_id, memory_type, key, length=SHORT_ID_LENGTH) + + +def _slugify(value: Any, *, fallback: str = "memory") -> str: + return _shared_slugify(value, fallback=fallback) + + +def _clamp_confidence(value: Any) -> float: + try: + score = float(value) + except (TypeError, ValueError): + return 0.0 + return round(max(0.0, min(score, 1.0)), 4) + + +def _normalise_memory_type(value: Any) -> MemoryType | None: + if isinstance(value, MemoryType): + return value + try: + return MemoryType(str(value).strip().lower()) + except ValueError: + return None + + +def _redact_text(text: str) -> tuple[str, bool]: + redacted = _SECRET_VALUE_RE.sub("[REDACTED_SECRET]", text) + redacted = _EMAIL_RE.sub("[REDACTED_EMAIL]", redacted) + redacted = _PHONE_RE.sub("[REDACTED_PHONE]", redacted) + return redacted, redacted != text + + +def _sanitize_for_prompt(value: Any) -> Any: + """Remove sensitive values before event data is sent to the LLM.""" + + if isinstance(value, dict): + sanitized: dict[str, Any] = {} + for key, item in value.items(): + if _SECRET_KEY_RE.search(str(key)): + sanitized[str(key)] = "[REDACTED_SECRET]" + else: + sanitized[str(key)] = _sanitize_for_prompt(item) + return sanitized + if isinstance(value, list): + return [_sanitize_for_prompt(item) for item in value] + if isinstance(value, tuple): + return [_sanitize_for_prompt(item) for item in value] + if isinstance(value, set): + return [_sanitize_for_prompt(item) for item in sorted(value, key=lambda item: str(item))] + if isinstance(value, str): + return _redact_text(value)[0] + return value + + +def _event_payload(event: MemoryEvent) -> dict[str, Any]: + return { + "event_id": event.event_id, + "user_id": event.user_id, + "session_id": event.session_id, + "task_id": event.task_id, + "event_type": event.event_type.value, + "scenario": event.scenario.value, + "source": event.source, + "actor": event.actor, + "content": _sanitize_for_prompt(event.content), + "tool_name": event.tool_name, + "input": _sanitize_for_prompt(event.input), + "output": _sanitize_for_prompt(event.output), + "success": event.success, + "metadata": _sanitize_for_prompt(event.metadata), + "timestamp": event.timestamp.isoformat(), + } + + +def event_payload_for_demo(event: MemoryEvent) -> dict[str, Any]: + """Expose sanitized event payload for demos and review artifacts.""" + + return _event_payload(event) + + +def _has_extractable_payload(event: MemoryEvent) -> bool: + payload = _event_payload(event) + return any( + payload.get(field) not in (None, "", {}, []) + for field in ("content", "tool_name", "input", "output", "metadata") + ) + + +def build_memory_extraction_prompt( + events: list[MemoryEvent], + *, + rule_candidates: list[MemoryCandidate] | None = None, + mode: str = "event", +) -> str: + """Build the LLM extraction request. + + ``rule_candidates`` is accepted only for backward compatibility with older + call sites. It is intentionally not included in the prompt because B-side + extraction is now LLM-only rather than rule-based. + """ + + payload = { + "mode": mode, + "events": [_event_payload(event) for event in events], + "output_contract": { + "shape": "Return JSON with a candidates array.", + "candidate_target": "Each candidate must be directly reusable as a MemoryCandidate.", + "empty_case": "If no reusable long-term memory exists, return {'candidates': []}.", + }, + "extraction_policy": [ + "Use semantic understanding to decide whether information is long-term memory.", + "Do not depend on keyword lists, regex templates, text length thresholds, or rule candidates.", + "Extract preference, knowledge, workflow, template, tool, environment, safety, profile, task_state, or session_summary memory when supported by evidence.", + "Mark temporary/current-task-only information as is_memory_worthy=false or is_long_term=false.", + "Never output credentials or raw secrets. Use redacted evidence when needed.", + "Keep evidence and reason explicit so reviewers can understand why the memory exists.", + ], + } + return json.dumps(payload, ensure_ascii=False, sort_keys=True) + + +def build_workflow_boundary_prompt(events: list[MemoryEvent]) -> str: + payload = { + "mode": "workflow_boundary_detection", + "events": [_event_payload(event) for event in events], + "output_contract": { + "shape": "Return JSON with a boundaries array.", + "indexing": "start and end are zero-based inclusive indices in the supplied events array.", + "empty_case": "If no workflow boundary is present, return {'boundaries': []}.", + }, + "extraction_policy": [ + "Use semantic understanding of task phases and tool dependencies.", + "Do not use fixed keyword markers or hardcoded step patterns.", + "Return only boundaries supported by event evidence.", + ], + } + return json.dumps(payload, ensure_ascii=False, sort_keys=True) + + +def should_call_llm(event: MemoryEvent, rule_candidates: list[MemoryCandidate] | None = None) -> bool: + """Return whether an event contains any payload worth sending to an LLM. + + The decision no longer contains event-type shortcuts, text-length + thresholds, confidence thresholds, or regex signal checks. If there is + observable payload, the LLM is the extractor. + """ + + return _has_extractable_payload(event) + + +class CandidateValidator: + """Validate and sanitize LLM-produced candidates.""" + + def validate(self, candidate: MemoryCandidate) -> MemoryCandidate | None: + content = candidate.content.strip() + if not content: + return None + + metadata = copy.deepcopy(candidate.metadata) + if metadata.get("is_long_term") is False or metadata.get("is_memory_worthy") is False: + return None + + redacted_content, redacted = _redact_text(content) + contains_secret_key = _SECRET_KEY_RE.search(content) is not None + if contains_secret_key and candidate.memory_type is not MemoryType.SAFETY: + return None + + if redacted: + metadata["sensitive_redacted"] = True + + return replace( + candidate, + content=redacted_content, + confidence=_clamp_confidence(candidate.confidence), + metadata=metadata, + ) + + def validate_many(self, candidates: list[MemoryCandidate]) -> list[MemoryCandidate]: + valid: list[MemoryCandidate] = [] + for candidate in candidates: + checked = self.validate(candidate) + if checked is not None: + valid.append(checked) + return valid + + +class CandidateMerger: + """Merge LLM candidates from multiple passes or scopes.""" + + @staticmethod + def merge( + rule_candidates: list[MemoryCandidate], + llm_candidates: list[MemoryCandidate], + ) -> list[MemoryCandidate]: + # The parameter name ``rule_candidates`` is kept for compatibility. + # In the LLM-only design it means "existing candidates". + validator = CandidateValidator() + all_candidates = validator.validate_many(rule_candidates) + validator.validate_many(llm_candidates) + by_identity: dict[tuple[str, MemoryType, str, Scene], MemoryCandidate] = {} + + for candidate in all_candidates: + identity = (candidate.user_id, candidate.memory_type, candidate.key, candidate.scenario) + current = by_identity.get(identity) + if current is None or candidate.confidence > current.confidence: + by_identity[identity] = CandidateMerger._clone_candidate(candidate) + continue + if candidate.confidence == current.confidence: + by_identity[identity] = CandidateMerger._merge_duplicate(current, candidate) + + merged = list(by_identity.values()) + CandidateMerger._annotate_conflicts(merged) + merged.sort(key=lambda item: (item.memory_type.value, item.key, -item.confidence)) + return merged + + @staticmethod + def _clone_candidate(candidate: MemoryCandidate) -> MemoryCandidate: + return replace( + candidate, + source_events=list(dict.fromkeys(candidate.source_events)), + source_summaries=list(dict.fromkeys(candidate.source_summaries)), + tags=list(dict.fromkeys(candidate.tags)), + metadata=copy.deepcopy(candidate.metadata), + ) + + @staticmethod + def _merge_duplicate(left: MemoryCandidate, right: MemoryCandidate) -> MemoryCandidate: + metadata = copy.deepcopy(left.metadata) + metadata.update({k: v for k, v in right.metadata.items() if k not in metadata}) + metadata["merged_sources"] = sorted({left.source, right.source, *metadata.get("merged_sources", [])}) + return replace( + left, + source_events=list(dict.fromkeys([*left.source_events, *right.source_events])), + source_summaries=list(dict.fromkeys([*left.source_summaries, *right.source_summaries])), + tags=list(dict.fromkeys([*left.tags, *right.tags])), + metadata=metadata, + ) + + @staticmethod + def _category_key(candidate: MemoryCandidate) -> str: + parts = candidate.key.split(".") + if len(parts) >= 2: + return ".".join(parts[:2]) + return candidate.key + + @staticmethod + def _annotate_conflicts(candidates: list[MemoryCandidate]) -> None: + groups: dict[tuple[str, MemoryType, Scene, str], list[MemoryCandidate]] = {} + for candidate in candidates: + group_id = ( + candidate.user_id, + candidate.memory_type, + candidate.scenario, + CandidateMerger._category_key(candidate), + ) + groups.setdefault(group_id, []).append(candidate) + + for group in groups.values(): + if len(group) <= 1: + continue + keys = sorted({candidate.key for candidate in group}) + for candidate in group: + candidate.metadata["possible_conflict_keys"] = [key for key in keys if key != candidate.key] + + +class LLMMemoryExtractor: + """Semantic memory extractor backed by an injected LLM client.""" + + @staticmethod + def extract_event( + event: MemoryEvent, + llm_client: LLMJsonClient, + *, + rule_candidates: list[MemoryCandidate] | None = None, + mode: str = "event", + ) -> list[MemoryCandidate]: + return LLMMemoryExtractor.extract_events( + [event], + llm_client, + rule_candidates=rule_candidates, + mode=mode, + ) + + @staticmethod + def extract_events( + events: list[MemoryEvent], + llm_client: LLMJsonClient, + *, + rule_candidates: list[MemoryCandidate] | None = None, + mode: str = "events", + ) -> list[MemoryCandidate]: + if not events: + return [] + active_events = [event for event in events if _has_extractable_payload(event)] + if not active_events: + return [] + + prompt = build_memory_extraction_prompt(active_events, rule_candidates=rule_candidates, mode=mode) + raw = llm_client.complete_json(prompt, MEMORY_EXTRACTION_SCHEMA) + payloads = LLMMemoryExtractor._candidate_payloads(raw) + candidates: list[MemoryCandidate] = [] + for payload in payloads: + candidate = LLMMemoryExtractor._candidate_from_payload(payload, active_events) + if candidate is not None: + candidates.append(candidate) + validated = CandidateMerger.merge([], CandidateValidator().validate_many(candidates)) + logger.debug( + "LLMMemoryExtractor.extract_events mode=%s events=%d payloads=%d candidates=%d validated=%d", + mode, + len(active_events), + len(payloads), + len(candidates), + len(validated), + ) + return validated + + @staticmethod + def _candidate_payloads(raw: Any) -> list[dict[str, Any]]: + if isinstance(raw, list): + return [item for item in raw if isinstance(item, dict)] + if isinstance(raw, dict): + candidates = raw.get("candidates", []) + if isinstance(candidates, list): + return [item for item in candidates if isinstance(item, dict)] + return [] + + @staticmethod + def _candidate_from_payload(payload: dict[str, Any], events: list[MemoryEvent]) -> MemoryCandidate | None: + if payload.get("is_memory_worthy") is False or payload.get("is_long_term") is False: + return None + + memory_type = _normalise_memory_type(payload.get("memory_type")) + if memory_type is None: + return None + + anchor = events[0] + content = str(payload.get("content") or "").strip() + category = _slugify(payload.get("category") or "general") + value = _slugify(payload.get("value") or "unspecified") + if not content: + return None + + key = f"{memory_type.value}.{category}.{value}" + extra_metadata = payload.get("metadata") if isinstance(payload.get("metadata"), dict) else {} + metadata = { + "extraction_method": "llm_semantic", + "schema_version": 2, + "is_memory_worthy": payload.get("is_memory_worthy", True), + "is_long_term": payload.get("is_long_term", True), + "category": category, + "value": str(payload.get("value") or ""), + "scope": str(payload.get("scope") or ""), + "evidence": str(payload.get("evidence") or ""), + "reason": str(payload.get("reason") or ""), + "sensitivity": str(payload.get("sensitivity") or "unknown"), + "conflicts_with": payload.get("conflicts_with") if isinstance(payload.get("conflicts_with"), list) else [], + "source_event_ids": [event.event_id for event in events], + "provenance": [ + { + "event_id": event.event_id, + "event_type": event.event_type.value, + "timestamp": event.timestamp.isoformat(), + } + for event in events + ], + } + metadata.update(extra_metadata) + + return MemoryCandidate( + candidate_id=_stable_id(anchor.user_id, memory_type, key), + user_id=anchor.user_id, + memory_type=memory_type, + key=key, + content=content, + scenario=anchor.scenario, + confidence=_clamp_confidence(payload.get("confidence")), + source="llm_extracted", + source_events=[event.event_id for event in events], + source_summaries=[str(payload.get("evidence") or "")], + tags=list(dict.fromkeys(["llm", memory_type.value, category])), + metadata=metadata, + ) + + +class LLMWorkflowBoundaryExtractor: + """LLM-backed workflow boundary detection.""" + + @staticmethod + def detect_boundaries(events: list[MemoryEvent], llm_client: LLMJsonClient) -> list[tuple[int, int]]: + if not events: + return [] + prompt = build_workflow_boundary_prompt(events) + raw = llm_client.complete_json(prompt, WORKFLOW_BOUNDARY_SCHEMA) + boundaries = raw.get("boundaries", []) if isinstance(raw, dict) else [] + parsed: list[tuple[int, int]] = [] + for item in boundaries: + if not isinstance(item, dict): + continue + try: + start = int(item["start"]) + end = int(item["end"]) + except (KeyError, TypeError, ValueError): + continue + if 0 <= start <= end < len(events): + parsed.append((start, end)) + logger.debug("LLMWorkflowBoundaryExtractor.detect_boundaries events=%d boundaries=%d", len(events), len(parsed)) + return parsed + + +def extract_candidates_with_default_client( + events: Iterable[MemoryEvent], + *, + mode: str, + memory_types: set[MemoryType] | None = None, +) -> list[MemoryCandidate]: + """Run the configured LLM client and optionally filter memory types.""" + + client = get_default_llm_client() + if client is None: + return [] + candidates = LLMMemoryExtractor.extract_events(list(events), client, mode=mode) + if memory_types is None: + return candidates + return [candidate for candidate in candidates if candidate.memory_type in memory_types] + + +def detect_boundaries_with_default_client(events: list[MemoryEvent]) -> list[tuple[int, int]]: + client = get_default_llm_client() + if client is None: + return [] + return LLMWorkflowBoundaryExtractor.detect_boundaries(events, client) + + +class HybridMemoryExtractor: + """Compatibility name for the LLM-only B-side extractor facade.""" + + @staticmethod + def extract_from_conversation( + event: MemoryEvent, + llm_client: LLMJsonClient | None = None, + ) -> list[MemoryCandidate]: + client = llm_client or get_default_llm_client() + if client is None: + return [] + result = LLMMemoryExtractor.extract_event(event, client, mode="conversation") + logger.debug("HybridMemoryExtractor.extract_from_conversation event_id=%s candidates=%d", event.event_id, len(result)) + return result + + @staticmethod + def extract_from_tool_result( + event: MemoryEvent, + llm_client: LLMJsonClient | None = None, + ) -> list[MemoryCandidate]: + client = llm_client or get_default_llm_client() + if client is None: + return [] + result = LLMMemoryExtractor.extract_event(event, client, mode="tool_result") + logger.debug("HybridMemoryExtractor.extract_from_tool_result event_id=%s candidates=%d", event.event_id, len(result)) + return result + + @staticmethod + def extract_from_session( + events: list[MemoryEvent], + llm_client: LLMJsonClient | None = None, + ) -> list[MemoryCandidate]: + client = llm_client or get_default_llm_client() + if client is None: + return [] + result = LLMMemoryExtractor.extract_events(events, client, mode="session") + logger.debug("HybridMemoryExtractor.extract_from_session events=%d candidates=%d", len(events), len(result)) + return result diff --git a/extractors/preference_extractor.py b/extractors/preference_extractor.py index c3a5bc8..c548ffb 100644 --- a/extractors/preference_extractor.py +++ b/extractors/preference_extractor.py @@ -1,681 +1,69 @@ +"""Preference extraction facade backed entirely by the configured LLM client. + +Public method names are kept for Phase 1 compatibility. The implementation no +longer contains regex templates, keyword lists, frequency rules, or length +thresholds. Callers that still use these legacy signatures must configure a +default LLM client via ``set_default_llm_client`` from +``extractors.llm_memory_extractor``. +""" + from __future__ import annotations -import hashlib -import json -import re -from collections import defaultdict -from dataclasses import replace -from typing import Any, Iterable +from datetime import datetime, timezone from core.constants import EventType, MemoryType, Scene from core.models import MemoryCandidate, MemoryEvent +from .llm_memory_extractor import extract_candidates_with_default_client -_FORMAT_ALIASES = { - "md": "markdown", - "markdown": "markdown", - "pdf": "pdf", - "word": "word", - "doc": "word", - "docx": "word", - "excel": "excel", - "xls": "excel", - "xlsx": "excel", - "csv": "csv", - "json": "json", - "jsonl": "jsonl", - "ppt": "powerpoint", - "pptx": "powerpoint", -} - -_STYLE_ALIASES = { - "简洁": "concise", - "详细": "detailed", - "正式": "formal", - "口语化": "conversational", - "分点": "bullets", - "结构化": "structured", - "先结论后分析": "conclusion_first", -} - -_LANGUAGE_ALIASES = { - "中文": "zh", - "英文": "en", - "中英双语": "bilingual", - "chinese": "zh", - "english": "en", -} - -_TOOL_ALIASES = { - "python": "python", - "bash": "bash", - "powershell": "powershell", - "git": "git", - "sqlite": "sqlite", - "curl": "curl", - "rg": "rg", - "ripgrep": "rg", -} - -_MAX_EXPLICIT_TEXT_CHARS = 12000 -_FORMAT_VALUE_RE = re.compile( - r"(?Markdown|MD|PDF|Word|DOCX|Excel|XLSX|CSV|JSONL|JSON|PPTX)(?![0-9A-Za-z])", - re.IGNORECASE, -) -_LANGUAGE_VALUE_RE = re.compile( - r"(?P中文|英文|中英双语|English|Chinese)", - re.IGNORECASE, -) -_STYLE_VALUE_RE = re.compile(r"(?P简洁|详细|正式|口语化|分点|结构化|先结论后分析)") -_PREFERENCE_CUE_RE = re.compile( - r"(以后|今后|下次|之后|默认|每次|始终|统一|偏好|喜欢|更喜欢|习惯|倾向|请|麻烦|希望|导出|保存|生成|输出|回复|回答|格式|用|使用|采用)", - re.IGNORECASE, -) -_CLAUSE_SPLIT_RE = re.compile(r"[\n\r。!?;;.!?]+") - -_PREFERENCE_PATTERNS: list[dict[str, Any]] = [ - { - "key": "output_format", - "pattern": re.compile( - r"(?:以后|今后|下次|之后)(?:都|请|麻烦|默认)?(?:用|使用|采用)\s*(?PMarkdown|MD|PDF|Word|DOCX|Excel|XLSX|CSV|JSONL|JSON|PPTX)", - re.IGNORECASE, - ), - "summary": "偏好以后使用 {value}", - }, - { - "key": "output_format", - "pattern": re.compile( - r"(?:导出|保存|生成)(?:为|成)?\s*(?PMarkdown|MD|PDF|Word|DOCX|Excel|XLSX|CSV|JSONL|JSON|PPTX)", - re.IGNORECASE, - ), - "summary": "偏好输出为 {value}", - }, - { - "key": "response_style", - "pattern": re.compile( - r"(?:我|用户)?(?:喜欢|偏好|更喜欢|习惯于|倾向于)\s*(?P简洁|详细|正式|口语化|分点|结构化|先结论后分析)", - ), - "summary": "偏好回答风格为 {value}", - }, - { - "key": "output_format", - "pattern": re.compile( - r"(?:我|用户)?(?:喜欢|偏好|更喜欢|习惯于|倾向于)\s*(?PMarkdown|MD|PDF|Word|DOCX|Excel|XLSX|CSV|JSONL|JSON|PPTX)", - re.IGNORECASE, - ), - "summary": "偏好输出为 {value}", - }, - { - "key": "language", - "pattern": re.compile( - r"(?:我|用户)?(?:喜欢|偏好|更喜欢|习惯于|倾向于)\s*(?P中文|英文|中英双语|English|Chinese)", - re.IGNORECASE, - ), - "summary": "偏好使用 {value}", - }, - { - "key": "tool", - "pattern": re.compile( - r"(?:我|用户)?(?:喜欢|偏好|更喜欢|习惯于|倾向于)\s*(?Ppython|bash|powershell|git|sqlite|curl|rg|ripgrep)", - re.IGNORECASE, - ), - "summary": "偏好优先使用 {value}", - }, - { - "key": "avoidance", - "pattern": re.compile( - r"(?:我|用户)?(?:不喜欢|不要|不想|避免|别)\s*(?P冗长|啰嗦|表格|代码块|废话|推测)", - ), - "summary": "偏好避免 {value}", - }, - { - "key": "language", - "pattern": re.compile( - r"(?:以后|今后|下次)(?:都|请|默认)?(?:用|使用|说|回复|回答)\s*(?P中文|英文|中英双语|English|Chinese)", - re.IGNORECASE, - ), - "summary": "偏好使用 {value}", - }, - { - "key": "tool", - "pattern": re.compile( - r"(?:请|麻烦)?(?:优先|尽量|默认)(?:使用|用)\s*(?Ppython|bash|powershell|git|sqlite|curl|rg|ripgrep)", - re.IGNORECASE, - ), - "summary": "偏好优先使用 {value}", - }, - { - "key": "workflow", - "pattern": re.compile( - r"(?:以后|今后|下次)(?:都|请|默认)?(?P先给结论后分析|先给结论|先列清单|先说风险|先给方案)", - ), - "summary": "偏好工作流 {value}", - }, - { - "key": "response_length", - "pattern": re.compile( - r"(?:回答|回复)(?:请|尽量)?(?P简短|简洁|详细|长一点|少一点)", - ), - "summary": "偏好回复长度 {value}", - }, - { - "key": "presentation", - "pattern": re.compile( - r"(?:请|麻烦)?(?:始终|一律|统一|每次都)?(?P分点|结构化|表格式|清单式|逐步)", - ), - "summary": "偏好表达方式 {value}", - }, - { - "key": "behavior", - "pattern": re.compile( - r"(?:每次|以后|今后|默认)(?:都|请)?(?:先|优先)?(?P确认|给草稿|给方案|给清单|给结论)", - ), - "summary": "偏好行为 {value}", - }, - { - "key": "language", - "pattern": re.compile( - r"(?:请|麻烦)?(?:全部|统一|默认)?(?:用|使用|写成)\s*(?P中文|英文|中英双语)", - ), - "summary": "偏好语言 {value}", - }, -] - - -def _slugify(value: str) -> str: - token = re.sub(r"[^0-9a-zA-Z\u4e00-\u9fff]+", "_", value.strip().lower()) - token = token.strip("_") - return token or "preference" - - -def _stable_candidate_id(user_id: str, memory_type: MemoryType, key: str) -> str: - payload = f"{user_id}\x1f{memory_type.value}\x1f{key}".encode("utf-8") - return hashlib.sha256(payload).hexdigest()[:32] - - -def _normalize_value(key: str, value: str) -> tuple[str, str]: - raw_value = value.strip() - if key == "output_format": - normalized = _FORMAT_ALIASES.get(raw_value.lower(), raw_value.lower()) - return normalized, normalized.upper() if normalized != "markdown" else "Markdown" - if key == "response_style": - normalized = _STYLE_ALIASES.get(raw_value, raw_value.lower()) - return normalized, raw_value - if key == "language": - normalized = _LANGUAGE_ALIASES.get(raw_value.lower(), _LANGUAGE_ALIASES.get(raw_value, raw_value.lower())) - return normalized, raw_value - if key == "tool": - normalized = _TOOL_ALIASES.get(raw_value.lower(), raw_value.lower()) - return normalized, raw_value - if key in {"workflow", "response_length", "presentation", "behavior", "avoidance"}: - return _slugify(raw_value), raw_value - return _slugify(raw_value), raw_value - - -def _canonical_structured_key(key: str) -> str: - normalized = key.replace("preferred_", "").replace("default_", "") - if "format" in normalized: - return "output_format" - if "language" in normalized or "locale" in normalized: - return "language" - if "style" in normalized or "tone" in normalized: - return "response_style" - if "tool" in normalized: - return "tool" - if "length" in normalized: - return "response_length" - if "workflow" in normalized or "mode" in normalized: - return "workflow" - if "parameter" in normalized or "option" in normalized: - return "parameter" - return normalized - -def _candidate_for( - *, - user_id: str, - scenario: Scene | str, - key: str, - normalized_value: str, - display_value: str, - content: str, - source: str, - source_events: list[str] | None = None, - source_summaries: list[str] | None = None, - tags: list[str] | None = None, - confidence: float = 0.9, - metadata: dict[str, Any] | None = None, -) -> MemoryCandidate: - candidate = MemoryCandidate( - candidate_id=_stable_candidate_id(user_id, MemoryType.PREFERENCE, f"preference.{key}.{normalized_value}"), - user_id=user_id, - memory_type=MemoryType.PREFERENCE, - key=f"preference.{key}.{normalized_value}", +def _synthetic_conversation_event(content: str) -> MemoryEvent: + return MemoryEvent( + event_id="synthetic-preference-content", + raw_event_id="synthetic-preference-content", + user_id="", + session_id="", + task_id="", + event_type=EventType.CONVERSATION, + scenario=Scene.GLOBAL, + source="conversation", + actor="user", content=content, - scenario=scenario, - confidence=confidence, - source=source, - source_events=source_events or [], - source_summaries=source_summaries or [], - tags=tags or [], - metadata=metadata or {}, + timestamp=datetime.now(timezone.utc), ) - candidate.metadata.setdefault("normalized_value", normalized_value) - candidate.metadata.setdefault("display_value", display_value) - return candidate - - -def _iter_text_fragments(payload: Any) -> Iterable[str]: - if payload is None: - return [] - if isinstance(payload, str): - text = payload.strip() - return [text] if text else [] - if isinstance(payload, (int, float, bool)): - return [str(payload)] - if isinstance(payload, dict): - fragments: list[str] = [] - for key, value in payload.items(): - if key in {"event_id", "raw_event_id", "timestamp", "created_at", "updated_at"}: - continue - if isinstance(value, (str, int, float, bool)): - fragments.extend(_iter_text_fragments(value)) - elif isinstance(value, (dict, list, tuple)): - fragments.extend(_iter_text_fragments(value)) - return fragments - if isinstance(payload, (list, tuple, set)): - fragments: list[str] = [] - items = sorted(payload, key=lambda item: str(item)) if isinstance(payload, set) else payload - for item in items: - fragments.extend(_iter_text_fragments(item)) - return fragments - return [str(payload)] - - -def _structured_pref_candidates(event: MemoryEvent) -> list[MemoryCandidate]: - candidates: list[MemoryCandidate] = [] - payload_sources = [event.input, event.output, event.metadata] - seen: set[tuple[str, str]] = set() - - for payload in payload_sources: - if not isinstance(payload, dict): - continue - for raw_key, raw_value in payload.items(): - if raw_value is None: - continue - key = str(raw_key).strip().lower() - if not isinstance(raw_value, (str, int, float, bool)): - continue - value_text = str(raw_value).strip() - if not value_text: - continue - if not any( - token in key - for token in ( - "format", - "style", - "tone", - "language", - "locale", - "tool", - "parameter", - "preference", - "preference", - "mode", - "workflow", - "response", - "length", - ) - ): - continue - canonical_key = _canonical_structured_key(key) - normalized, display_value = _normalize_value(canonical_key, value_text) - pair = (canonical_key, normalized) - if pair in seen: - continue - seen.add(pair) - human_key = canonical_key.replace("_", " ") - content = f"偏好{human_key}为 {display_value}" - candidates.append( - _candidate_for( - user_id=event.user_id, - scenario=event.scenario, - key=canonical_key, - normalized_value=normalized, - display_value=display_value, - content=content, - source="tool_result", - source_events=[event.event_id], - source_summaries=[content], - tags=["explicit", "structured"], - confidence=0.86, - metadata={ - "match_type": "structured_field", - "field": canonical_key, - "field_value": value_text, - }, - ) - ) - return candidates - - -def _extract_from_text(content: str, *, source: str = "conversation") -> list[MemoryCandidate]: - text = content or "" - if not text.strip(): - return [] - if len(text) > _MAX_EXPLICIT_TEXT_CHARS: - text = text[:_MAX_EXPLICIT_TEXT_CHARS] - - results: list[MemoryCandidate] = [] - seen: set[tuple[str, str]] = set() - - for spec in _PREFERENCE_PATTERNS: - for match in spec["pattern"].finditer(text): - raw_value = match.group("value").strip() - if not raw_value: - continue - normalized_value, display_value = _normalize_value(spec["key"], raw_value) - pair = (spec["key"], normalized_value) - if pair in seen: - continue - seen.add(pair) - summary = spec["summary"].format(value=display_value) - confidence = 0.95 if source == "conversation" else 0.9 - results.append( - _candidate_for( - user_id="", - scenario=Scene.GLOBAL, - key=spec["key"], - normalized_value=normalized_value, - display_value=display_value, - content=summary, - source=source, - source_summaries=[text], - tags=["explicit", spec["key"]], - confidence=confidence, - metadata={ - "match_type": "pattern", - "pattern": spec["pattern"].pattern, - "matched_text": match.group(0), - }, - ) - ) - - for candidate in _extract_flexible_preferences(text, source=source): - pair = (candidate.metadata.get("field", ""), candidate.metadata.get("normalized_value", "")) - if pair in seen: - continue - seen.add(pair) - results.append(candidate) - - return results - - -def _extract_flexible_preferences(text: str, *, source: str) -> list[MemoryCandidate]: - """Extract natural-language preferences when modifiers appear between cues. - - Rule-based patterns above are intentionally precise. Real user utterances - often insert task words between the temporal cue and the value, e.g. - "以后导出都用 PDF 格式". This fallback keeps the extraction deterministic - while allowing those modifiers. - """ - results: list[MemoryCandidate] = [] - seen: set[tuple[str, str]] = set() - - for raw_clause in _CLAUSE_SPLIT_RE.split(text): - clause = raw_clause.strip() - if not clause or len(clause) > 240: - continue - if not _PREFERENCE_CUE_RE.search(clause): - continue - - for key, value_re, summary in ( - ("output_format", _FORMAT_VALUE_RE, "偏好输出为 {value}"), - ("language", _LANGUAGE_VALUE_RE, "偏好使用 {value}"), - ("response_style", _STYLE_VALUE_RE, "偏好回答风格为 {value}"), - ): - for match in value_re.finditer(clause): - raw_value = match.group("value").strip() - normalized_value, display_value = _normalize_value(key, raw_value) - pair = (key, normalized_value) - if pair in seen: - continue - seen.add(pair) - confidence = 0.9 if source == "conversation" else 0.84 - results.append( - _candidate_for( - user_id="", - scenario=Scene.GLOBAL, - key=key, - normalized_value=normalized_value, - display_value=display_value, - content=summary.format(value=display_value), - source=source, - source_summaries=[clause], - tags=["explicit", key, "flexible"], - confidence=confidence, - metadata={ - "match_type": "flexible_clause", - "field": key, - "matched_text": clause, - }, - ) - ) - return results - - -def _rebind_candidates( - candidates: list[MemoryCandidate], - *, - event: MemoryEvent, - source: str, -) -> list[MemoryCandidate]: - rebound: list[MemoryCandidate] = [] - for candidate in candidates: - rebound.append( - replace( - candidate, - candidate_id=_stable_candidate_id(event.user_id, candidate.memory_type, candidate.key), - user_id=event.user_id, - scenario=event.scenario, - source=source, - source_events=(candidate.source_events or []) + [event.event_id], - source_summaries=(candidate.source_summaries or []) or [event.content or event.source], - ) - ) - return rebound class PreferenceExtractor: @staticmethod def extract_from_conversation(event: MemoryEvent) -> list[MemoryCandidate]: - pieces = [event.content, *(_iter_text_fragments(event.input)), *(_iter_text_fragments(event.output))] - text = "\n".join(piece for piece in pieces if piece) - candidates = _extract_from_text(text, source="conversation") - candidates.extend(_structured_pref_candidates(event)) - return _dedupe_candidates(_rebind_candidates(candidates, event=event, source="conversation")) + return extract_candidates_with_default_client( + [event], + mode="preference_from_conversation", + memory_types={MemoryType.PREFERENCE}, + ) @staticmethod def extract_from_tool_result(event: MemoryEvent) -> list[MemoryCandidate]: - pieces = [ - event.content, - *(_iter_text_fragments(event.output)), - *(_iter_text_fragments(event.metadata)), - ] - text = "\n".join(piece for piece in pieces if piece) - candidates = _extract_from_text(text, source="tool_result") - candidates.extend(_structured_pref_candidates(event)) - return _dedupe_candidates(_rebind_candidates(candidates, event=event, source="tool_result")) + return extract_candidates_with_default_client( + [event], + mode="preference_from_tool_result", + memory_types={MemoryType.PREFERENCE}, + ) @staticmethod def extract_explicit_preference(content: str) -> list[MemoryCandidate]: - return _extract_from_text(content, source="conversation") + if not content or not content.strip(): + return [] + return extract_candidates_with_default_client( + [_synthetic_conversation_event(content)], + mode="preference_from_text", + memory_types={MemoryType.PREFERENCE}, + ) @staticmethod def extract_implicit_preference(events: list[MemoryEvent]) -> list[MemoryCandidate]: - by_user: dict[str, list[MemoryEvent]] = defaultdict(list) - for event in events: - by_user[event.user_id].append(event) - - candidates: list[MemoryCandidate] = [] - for user_id, user_events in by_user.items(): - candidates.extend(_extract_implicit_for_user(user_id, user_events)) - return _dedupe_candidates(candidates) - - -def _dedupe_candidates(candidates: list[MemoryCandidate]) -> list[MemoryCandidate]: - unique: list[MemoryCandidate] = [] - seen: set[tuple[str, str, str]] = set() - for candidate in candidates: - signature = (candidate.user_id, candidate.memory_type.value, candidate.key) - if signature in seen: - continue - seen.add(signature) - unique.append(candidate) - return unique - - -def _normalize_scalar(value: Any) -> str: - if isinstance(value, bool): - return "true" if value else "false" - if isinstance(value, (int, float)): - return str(value) - return str(value).strip() - - -def _collect_scalar_params(payload: Any, prefix: str = "") -> list[tuple[str, str]]: - collected: list[tuple[str, str]] = [] - if not isinstance(payload, dict): - return collected - for raw_key, raw_value in payload.items(): - key = f"{prefix}{raw_key}".strip(".") - if raw_value is None: - continue - if isinstance(raw_value, dict): - collected.extend(_collect_scalar_params(raw_value, prefix=f"{key}.")) - elif isinstance(raw_value, (list, tuple, set)): - values = sorted(raw_value, key=lambda item: str(item)) if isinstance(raw_value, set) else raw_value - joined = ", ".join(_normalize_scalar(item) for item in values if item is not None) - if joined: - collected.append((key.lower(), joined)) - elif isinstance(raw_value, (str, int, float, bool)): - text = _normalize_scalar(raw_value) - if text: - collected.append((key.lower(), text)) - return collected - - -def _action_signature(event: MemoryEvent) -> str: - if event.content and event.content.strip(): - return re.sub(r"\s+", " ", event.content.strip().lower()) - if event.tool_name: - return f"tool:{event.tool_name.strip().lower()}" - if isinstance(event.metadata, dict): - for key in ("action", "intent", "behavior", "operation"): - value = event.metadata.get(key) - if isinstance(value, str) and value.strip(): - return f"{key}:{re.sub(r'\s+', ' ', value.strip().lower())}" - return f"{event.event_type.value}:{event.source.strip().lower()}" - - -def _extract_implicit_for_user(user_id: str, events: list[MemoryEvent]) -> list[MemoryCandidate]: - if not events: - return [] - - ordered_events = sorted(events, key=lambda event: event.timestamp) - results: list[MemoryCandidate] = [] - - action_groups: dict[str, list[MemoryEvent]] = defaultdict(list) - tool_groups: dict[str, list[MemoryEvent]] = defaultdict(list) - parameter_groups: dict[tuple[str, str], list[MemoryEvent]] = defaultdict(list) - - for event in ordered_events: - action_groups[_action_signature(event)].append(event) - if event.tool_name: - tool_groups[event.tool_name.strip().lower()].append(event) - for key, value in _collect_scalar_params(event.input): - if any(token in key for token in ("format", "language", "style", "mode", "length", "temperature", "top_p", "prompt", "model", "tool")): - parameter_groups[(key, value)].append(event) - - for signature, group in action_groups.items(): - if len(group) < 3: - continue - results.append( - _candidate_for( - user_id=user_id, - scenario=_shared_scenario(group), - key="workflow." + _slugify(signature), - normalized_value=_slugify(signature), - display_value=signature, - content=f"用户连续 {len(group)} 次重复相同操作:{group[0].content or group[0].source}", - source="implicit_action", - source_events=[event.event_id for event in group], - source_summaries=[event.content or event.source for event in group], - tags=["implicit", "workflow"], - confidence=min(0.55 + 0.1 * (len(group) - 3), 0.95), - metadata={ - "heuristic": "repeated_action", - "count": len(group), - "signature": signature, - }, - ) + return extract_candidates_with_default_client( + events, + mode="preference_from_session", + memory_types={MemoryType.PREFERENCE}, ) - - for tool_name, group in tool_groups.items(): - if len(group) < 3: - continue - display_tool = tool_name - results.append( - _candidate_for( - user_id=user_id, - scenario=_shared_scenario(group), - key="tool." + _slugify(tool_name), - normalized_value=_slugify(tool_name), - display_value=display_tool, - content=f"用户连续 {len(group)} 次使用工具 {display_tool}", - source="implicit_tool", - source_events=[event.event_id for event in group], - source_summaries=[event.content or event.tool_name or event.source for event in group], - tags=["implicit", "tool"], - confidence=min(0.55 + 0.1 * (len(group) - 3), 0.95), - metadata={ - "heuristic": "repeated_tool", - "count": len(group), - "tool_name": tool_name, - }, - ) - ) - - for (param_name, param_value), group in parameter_groups.items(): - if len(group) < 3: - continue - results.append( - _candidate_for( - user_id=user_id, - scenario=_shared_scenario(group), - key=f"parameter.{_slugify(param_name)}", - normalized_value=_slugify(param_value), - display_value=param_value, - content=f"用户多次使用参数 {param_name}={param_value}", - source="implicit_parameter", - source_events=[event.event_id for event in group], - source_summaries=[event.content or event.source for event in group], - tags=["implicit", "parameter"], - confidence=min(0.55 + 0.1 * (len(group) - 3), 0.95), - metadata={ - "heuristic": "repeated_parameter", - "count": len(group), - "parameter_name": param_name, - "parameter_value": param_value, - }, - ) - ) - - return results - - -def _shared_scenario(events: list[MemoryEvent]) -> Scene: - scenario = events[0].scenario - if all(event.scenario == scenario for event in events): - return scenario - return Scene.GLOBAL diff --git a/extractors/tool_extractor.py b/extractors/tool_extractor.py index 19c34c9..2122542 100644 --- a/extractors/tool_extractor.py +++ b/extractors/tool_extractor.py @@ -1,323 +1,40 @@ -from __future__ import annotations - -import hashlib -import json -import re -from collections import Counter, defaultdict, deque -from dataclasses import dataclass -from typing import Any - -from core.constants import EventType, MemoryType, Scene -from core.models import MemoryCandidate, MemoryEvent - - -_CORRELATION_KEYS = ("tool_call_id", "call_id", "invocation_id", "request_id", "trace_id", "run_id") -_DURATION_KEYS = ("duration_ms", "latency_ms", "elapsed_ms", "response_time_ms", "execution_time_ms") -_SENSITIVE_TOKENS = ("password", "passwd", "secret", "token", "api_key", "access_key", "credential", "authorization") -_SUCCESS_STATUSES = {"ok", "success", "succeeded", "completed", "complete", "done", "passed"} -_FAILURE_STATUSES = {"error", "failed", "failure", "exception", "timeout", "timed_out", "cancelled", "canceled"} - - -@dataclass(frozen=True, slots=True) -class _ToolInvocation: - user_id: str - tool_name: str - scenario: Scene - event_ids: tuple[str, ...] - success: bool | None - duration_ms: float | None - failure_reason: str | None - parameters: tuple[tuple[str, str], ...] - - -def _slugify(value: str) -> str: - token = re.sub(r"[^0-9a-zA-Z\u4e00-\u9fff]+", "_", value.strip().lower()) - return token.strip("_") or "tool" - - -def _stable_candidate_id(user_id: str, memory_type: MemoryType, key: str) -> str: - payload = f"{user_id}\x1f{memory_type.value}\x1f{key}".encode("utf-8") - return hashlib.sha256(payload).hexdigest()[:32] - - -def _normalise_tool_name(event: MemoryEvent) -> str: - return (event.tool_name or "unknown_tool").strip().lower() or "unknown_tool" - - -def _is_tool_event(event: MemoryEvent) -> bool: - return event.event_type in {EventType.TOOL_CALL, EventType.TOOL_RESULT} or bool(event.tool_name) - - -def _group_key(event: MemoryEvent) -> tuple[str, str, str, str]: - return event.user_id, event.session_id, event.task_id, _normalise_tool_name(event) - - -def _correlation_id(event: MemoryEvent) -> str | None: - for payload in (event.metadata, event.output, event.input, event.raw_event or {}): - if not isinstance(payload, dict): - continue - for key in _CORRELATION_KEYS: - value = payload.get(key) - if value is not None and str(value).strip(): - return f"{key}:{str(value).strip()}" - return None - - -def _as_duration(value: Any) -> float | None: - if isinstance(value, bool): - return None - if isinstance(value, (int, float)) and value >= 0: - return float(value) - if isinstance(value, str): - try: - parsed = float(value.strip()) - except ValueError: - return None - return parsed if parsed >= 0 else None - return None - - -def _duration_ms(event: MemoryEvent) -> float | None: - for payload in (event.metadata, event.output, event.input): - if not isinstance(payload, dict): - continue - for key in _DURATION_KEYS: - duration = _as_duration(payload.get(key)) - if duration is not None: - return duration - return None - - -def _terminal_success(event: MemoryEvent) -> bool | None: - if event.success is not None: - return bool(event.success) - - text_values: list[str] = [] - for payload in (event.output, event.metadata): - if not isinstance(payload, dict): - continue - for key in ("status", "state", "result", "outcome"): - value = payload.get(key) - if isinstance(value, str): - text_values.append(value.lower().strip()) - if any(key in payload and payload[key] not in (None, "", False) for key in ("error", "exception", "stderr", "error_type")): - return False - - text_values.extend(re.findall(r"[a-z_]+", (event.content or "").lower())) - if any(value in _FAILURE_STATUSES for value in text_values): - return False - if any(value in _SUCCESS_STATUSES for value in text_values): - return True - return None +"""Tool memory extraction facade backed entirely by the configured LLM. +The extractor no longer computes tool patterns with counters, hardcoded +success formulas, or parameter-frequency rules. Tool experience is extracted +from model-produced ``MemoryCandidate`` objects. The legacy success-rate +method reads a model-provided ``metadata.success_rate`` field when available. +""" -def _normalise_failure_reason(value: str) -> str: - text = re.sub(r"[A-Za-z]:\\[^\s,;]+|/[^\s,;]+", "", value.strip().lower()) - text = re.sub(r"\b\d+\b", "", text) - text = re.sub(r"\s+", " ", text) - return text[:160] or "unknown_failure" - - -def _failure_reason(event: MemoryEvent) -> str: - for payload in (event.output, event.metadata, event.raw_event or {}): - if not isinstance(payload, dict): - continue - for key in ("error_type", "error_code", "error", "exception", "stderr", "message", "reason"): - value = payload.get(key) - if isinstance(value, str) and value.strip(): - return _normalise_failure_reason(f"{key}: {value}") - if event.content and event.content.strip(): - return _normalise_failure_reason(event.content) - return "unknown_failure" - - -def _collect_parameters(payload: Any, prefix: str = "") -> list[tuple[str, str]]: - if not isinstance(payload, dict): - return [] - - parameters: list[tuple[str, str]] = [] - for raw_key, value in sorted(payload.items(), key=lambda item: str(item[0])): - key = f"{prefix}.{raw_key}".strip(".").lower() - if any(token in key for token in _SENSITIVE_TOKENS): - continue - if value is None: - continue - if isinstance(value, dict): - parameters.extend(_collect_parameters(value, key)) - continue - if isinstance(value, set): - value = sorted(value, key=str) - if isinstance(value, (list, tuple)): - rendered = json.dumps(value, ensure_ascii=False, sort_keys=True, default=str) - elif isinstance(value, (str, int, float, bool)): - rendered = str(value).strip() - else: - continue - if rendered: - parameters.append((key, rendered[:120])) - return parameters - - -def _event_scenario(events: list[MemoryEvent]) -> Scene: - scenario = events[0].scenario - return scenario if all(event.scenario == scenario for event in events) else Scene.SYSTEM - - -def _build_invocations(events: list[MemoryEvent]) -> list[_ToolInvocation]: - grouped: dict[tuple[str, str, str, str], list[MemoryEvent]] = defaultdict(list) - for event in events: - if _is_tool_event(event): - grouped[_group_key(event)].append(event) - - invocations: list[_ToolInvocation] = [] - for (_, _, _, tool_name), group in sorted(grouped.items()): - ordered = sorted(group, key=lambda event: (event.timestamp, event.event_id)) - pending_by_id: dict[str, deque[MemoryEvent]] = defaultdict(deque) - pending_fifo: deque[MemoryEvent] = deque() - - for event in ordered: - if event.event_type is EventType.TOOL_CALL: - pending_fifo.append(event) - correlation = _correlation_id(event) - if correlation: - pending_by_id[correlation].append(event) - continue - - if event.event_type is not EventType.TOOL_RESULT: - continue - - correlation = _correlation_id(event) - call_event: MemoryEvent | None = None - if correlation and pending_by_id[correlation]: - call_event = pending_by_id[correlation].popleft() - pending_fifo.remove(call_event) - elif pending_fifo: - call_event = pending_fifo.popleft() - call_correlation = _correlation_id(call_event) - if call_correlation and pending_by_id[call_correlation]: - pending_by_id[call_correlation].remove(call_event) - - duration = _duration_ms(event) - if duration is None and call_event is not None: - duration = _duration_ms(call_event) - if duration is None and call_event is not None: - elapsed = (event.timestamp - call_event.timestamp).total_seconds() * 1000 - duration = elapsed if elapsed >= 0 else None - - success = _terminal_success(event) - invocations.append( - _ToolInvocation( - user_id=event.user_id, - tool_name=tool_name, - scenario=_event_scenario([call_event, event] if call_event else [event]), - event_ids=tuple(event_id for event_id in ((call_event.event_id if call_event else None), event.event_id) if event_id), - success=success, - duration_ms=duration, - failure_reason=_failure_reason(event) if success is False else None, - parameters=tuple(_collect_parameters(call_event.input if call_event else event.input)), - ) - ) - - for call_event in pending_fifo: - invocations.append( - _ToolInvocation( - user_id=call_event.user_id, - tool_name=tool_name, - scenario=call_event.scenario, - event_ids=(call_event.event_id,), - success=None, - duration_ms=_duration_ms(call_event), - failure_reason=None, - parameters=tuple(_collect_parameters(call_event.input)), - ) - ) - return invocations +from __future__ import annotations +from core.constants import MemoryType +from core.models import MemoryCandidate, MemoryEvent -def _success_rate(invocations: list[_ToolInvocation]) -> float: - completed = [item for item in invocations if item.success is not None] - if not completed: - return 0.0 - return sum(item.success is True for item in completed) / len(completed) +from .llm_memory_extractor import extract_candidates_with_default_client class ToolExtractor: @staticmethod def calculate_tool_success_rate(tool_name: str, events: list[MemoryEvent]) -> float: - """Return success_count / completed_invocation_count for one tool. - - Calls without a terminal result are excluded rather than being silently - counted as failures; their number is reported by ``extract_tool_pattern``. - """ - normalized_tool = tool_name.strip().lower() - invocations = [item for item in _build_invocations(events) if item.tool_name == normalized_tool] - return _success_rate(invocations) + candidates = extract_candidates_with_default_client( + events, + mode=f"tool_success_rate:{tool_name}", + memory_types={MemoryType.TOOL}, + ) + for candidate in candidates: + if str(candidate.metadata.get("tool_name", "")).strip().lower() != tool_name.strip().lower(): + continue + try: + return max(0.0, min(float(candidate.metadata["success_rate"]), 1.0)) + except (KeyError, TypeError, ValueError): + continue + return 0.0 @staticmethod def extract_tool_pattern(events: list[MemoryEvent]) -> list[MemoryCandidate]: - by_user_tool: dict[tuple[str, str], list[_ToolInvocation]] = defaultdict(list) - for invocation in _build_invocations(events): - by_user_tool[(invocation.user_id, invocation.tool_name)].append(invocation) - - candidates: list[MemoryCandidate] = [] - for (user_id, tool_name), invocations in sorted(by_user_tool.items()): - completed = [item for item in invocations if item.success is not None] - successes = sum(item.success is True for item in completed) - failures = sum(item.success is False for item in completed) - unknown = len(invocations) - len(completed) - durations = [item.duration_ms for item in invocations if item.duration_ms is not None] - failure_counts = Counter(item.failure_reason for item in invocations if item.failure_reason) - parameter_counts = Counter(item.parameters for item in invocations if item.parameters) - rate = _success_rate(invocations) - completeness = len(completed) / len(invocations) if invocations else 0.0 - common_failures = [ - {"reason": reason, "count": count} - for reason, count in failure_counts.most_common(5) - ] - common_parameters = [ - {"parameters": dict(parameters), "count": count} - for parameters, count in parameter_counts.most_common(5) - ] - scenarios = [item.scenario for item in invocations] - scenario = scenarios[0] if all(item == scenarios[0] for item in scenarios) else Scene.SYSTEM - key = f"tool.pattern.{_slugify(tool_name)}" - content_parts = [ - f"工具 {tool_name}:成功率 {rate:.2%}({successes}/{len(completed)})", - f"数据完整度 {completeness:.2%}(已完成 {len(completed)},未完成 {unknown})", - ] - if durations: - content_parts.append(f"平均响应 {sum(durations) / len(durations):.2f}ms") - if common_failures: - content_parts.append(f"常见失败 {common_failures[0]['reason']}({common_failures[0]['count']})") - - event_ids = sorted({event_id for item in invocations for event_id in item.event_ids}) - candidates.append( - MemoryCandidate( - candidate_id=_stable_candidate_id(user_id, MemoryType.TOOL, key), - user_id=user_id, - memory_type=MemoryType.TOOL, - key=key, - content=";".join(content_parts), - scenario=scenario, - confidence=round(min(0.98, 0.55 + 0.45 * completeness), 4), - source="tool_pattern", - source_events=event_ids, - source_summaries=[f"{tool_name}: {item.success}" for item in invocations], - tags=["tool", "pattern", _slugify(tool_name)], - metadata={ - "tool_name": tool_name, - "success_count": successes, - "failure_count": failures, - "total_count": len(completed), - "observed_count": len(invocations), - "unknown_count": unknown, - "success_rate": round(rate, 4), - "data_completeness": round(completeness, 4), - "average_response_ms": round(sum(durations) / len(durations), 4) if durations else None, - "duration_sample_count": len(durations), - "common_failure_reasons": common_failures, - "common_parameter_combinations": common_parameters, - }, - ) - ) - return candidates + return extract_candidates_with_default_client( + events, + mode="tool_pattern", + memory_types={MemoryType.TOOL}, + ) diff --git a/extractors/workflow_extractor.py b/extractors/workflow_extractor.py index 0993fc5..bfff2ea 100644 --- a/extractors/workflow_extractor.py +++ b/extractors/workflow_extractor.py @@ -1,504 +1,35 @@ -from __future__ import annotations - -import hashlib -import re -from collections import defaultdict -from typing import Any, Iterable - -from core.constants import EventType, MemoryType, Scene -from core.models import MemoryCandidate, MemoryEvent - - -_WORKFLOW_TEXT_MARKERS = ( - "workflow", - "pipeline", - "sequence", - "step", - "steps", - "then", - "after", - "before", - "next", - "follow", - "run", - "execute", - "process", - "tool chain", - "dependency", - "dependent", - "flow", - "步骤", - "流程", - "然后", - "接着", - "再", - "最后", - "依赖", - "串联", - "串行", - "顺序", - "执行", - "运行", - "调用", - "工具", - "导出", - "合并", - "安装", - "配置", -) - -_TRANSITION_MARKERS = ( - "then", - "after", - "before", - "next", - "based on", - "using", - "with", - "result of", - "from", - "then use", - "然后", - "接着", - "再", - "之后", - "基于", - "使用", - "基于上一步", - "利用", - "再用", - "根据", -) - -_FILE_RE = re.compile( - r"\b[\w.-]+\.(?:csv|tsv|xlsx|xls|json|jsonl|txt|log|md|pdf|docx|zip|tar|gz|py|sh|sql)\b", - re.IGNORECASE, -) -_PATH_RE = re.compile(r"(?:[A-Za-z]:\\|/)[^\s,;]+") -_STEP_RE = re.compile(r"^(?:\d+[.)、-]|step\s*\d+|步骤\s*\d+)", re.IGNORECASE) - - -def _slugify(text: str) -> str: - token = re.sub(r"[^0-9a-zA-Z\u4e00-\u9fff]+", "_", text.strip().lower()) - token = re.sub(r"_+", "_", token).strip("_") - return token or "workflow" - - -def _stable_candidate_id(user_id: str, memory_type: MemoryType, key: str) -> str: - payload = f"{user_id}\x1f{memory_type.value}\x1f{key}".encode("utf-8") - return hashlib.sha256(payload).hexdigest()[:32] - - -def _normalize_text(value: Any) -> str: - text = str(value).strip().lower() - return re.sub(r"\s+", " ", text) - - -def _iter_text_fragments(payload: Any) -> list[str]: - if payload is None: - return [] - if isinstance(payload, str): - text = payload.strip() - return [text] if text else [] - if isinstance(payload, (int, float, bool)): - return [str(payload)] - if isinstance(payload, dict): - fragments: list[str] = [] - for key, value in payload.items(): - if key in {"event_id", "raw_event_id", "timestamp", "created_at", "updated_at"}: - continue - fragments.extend(_iter_text_fragments(value)) - return fragments - if isinstance(payload, (list, tuple, set)): - fragments: list[str] = [] - items = sorted(payload, key=lambda item: str(item)) if isinstance(payload, set) else payload - for item in items: - fragments.extend(_iter_text_fragments(item)) - return fragments - return [str(payload)] - - -def _event_text(event: MemoryEvent) -> str: - parts = [ - event.content or "", - event.tool_name or "", - *(_iter_text_fragments(event.input)), - *(_iter_text_fragments(event.output)), - *(_iter_text_fragments(event.metadata)), - ] - return "\n".join(part for part in parts if part and str(part).strip()) - - -def _is_tool_event(event: MemoryEvent) -> bool: - return event.event_type in {EventType.TOOL_CALL, EventType.TOOL_RESULT} or bool(event.tool_name) - - -def _is_workflow_relevant(event: MemoryEvent) -> bool: - if _is_tool_event(event): - return True - - text = _normalize_text(_event_text(event)) - if not text: - return False - - if any(marker in text for marker in _WORKFLOW_TEXT_MARKERS): - return True - - if _STEP_RE.search(text): - return True - - return False - - -def _extract_artifacts(event: MemoryEvent) -> set[str]: - artifacts: set[str] = set() - text = _event_text(event) - for match in _FILE_RE.finditer(text): - artifacts.add(match.group(0).lower()) - for match in _PATH_RE.finditer(text): - artifacts.add(match.group(0).lower()) - - for payload in (event.input, event.output, event.metadata): - if not isinstance(payload, dict): - continue - for key, value in payload.items(): - if value is None: - continue - key_text = str(key).strip().lower() - if key_text in {"source", "target", "file", "path", "input", "output", "depends_on", "previous", "next", "artifact", "artifacts"}: - for fragment in _iter_text_fragments(value): - fragment = fragment.strip().lower() - if fragment: - artifacts.add(fragment) - return artifacts - - -def _group_key(event: MemoryEvent) -> tuple[str, str, str]: - """Keep simultaneous tasks in one session isolated from each other.""" - return event.user_id, event.session_id, event.task_id - - -def _group_events(events: list[MemoryEvent]) -> list[tuple[tuple[str, str, str], list[MemoryEvent]]]: - grouped: dict[tuple[str, str, str], list[MemoryEvent]] = defaultdict(list) - for event in events: - key = _group_key(event) - grouped[key].append(event) - return [ - (key, sorted(group, key=lambda event: (event.timestamp, event.event_id))) - for key, group in sorted(grouped.items()) - ] - - -def _segment_scenario(events: list[MemoryEvent]) -> Scene: - scenario = events[0].scenario - if all(event.scenario == scenario for event in events): - return scenario - return Scene.GLOBAL - - -def _boundary_signature(events: list[MemoryEvent], boundaries: list[tuple[int, int]]) -> str: - parts: list[str] = [] - for start, end in boundaries: - segment = events[start : end + 1] - tool_names = [event.tool_name or event.event_type.value for event in segment if _is_tool_event(event)] - if tool_names: - parts.append("->".join(_slugify(name) for name in tool_names)) - return "__".join(parts) if parts else "workflow" - - -def _workflow_boundaries_with_indices(events: list[MemoryEvent]) -> list[tuple[int, int]]: - relevant = [index for index, event in enumerate(events) if _is_workflow_relevant(event)] - if not relevant: - return [] - - boundaries: list[tuple[int, int]] = [] - start = prev = relevant[0] - for index in relevant[1:]: - if index - prev <= 2: - prev = index - continue - boundaries.append((start, prev)) - start = prev = index - boundaries.append((start, prev)) - return boundaries - - -def _tool_groups(segment: list[MemoryEvent]) -> list[list[MemoryEvent]]: - groups: list[list[MemoryEvent]] = [] - for event in segment: - if not _is_tool_event(event): - continue - if groups and _same_tool_group(groups[-1][-1], event): - groups[-1].append(event) - continue - groups.append([event]) - return groups +"""Workflow extraction facade backed entirely by the configured LLM. +Workflow detection no longer uses tool-sequence heuristics, transition keyword +markers, file/path regexes, or reproduction-rate formulas. The configured LLM +must return either workflow candidates or explicit workflow boundaries. +""" -def _same_tool_group(left: MemoryEvent, right: MemoryEvent) -> bool: - left_tool = (left.tool_name or "").strip().lower() - right_tool = (right.tool_name or "").strip().lower() - if not left_tool or not right_tool: - return False - if left_tool != right_tool: - return False - if left.event_type == right.event_type: - return True - return {left.event_type, right.event_type} <= {EventType.TOOL_CALL, EventType.TOOL_RESULT} - - -def _step_label(group: list[MemoryEvent]) -> str: - representative = group[0] - if representative.tool_name: - return representative.tool_name.strip().lower() - if representative.content: - return re.sub(r"\s+", " ", representative.content.strip().lower())[:40] - return representative.event_type.value - - -def _group_artifacts(group: list[MemoryEvent]) -> set[str]: - artifacts: set[str] = set() - for event in group: - artifacts |= _extract_artifacts(event) - return artifacts - - -def _group_transition_markers(group: list[MemoryEvent]) -> bool: - text = " ".join(_normalize_text(_event_text(event)) for event in group) - return any(marker in text for marker in _TRANSITION_MARKERS) - - -def _dependency_edges(groups: list[list[MemoryEvent]]) -> list[dict[str, Any]]: - edges: list[dict[str, Any]] = [] - if len(groups) < 2: - return edges - - group_artifacts = [_group_artifacts(group) for group in groups] - group_texts = [" ".join(_normalize_text(_event_text(event)) for event in group) for group in groups] - - for index in range(len(groups) - 1): - left_group = groups[index] - right_group = groups[index + 1] - overlap = sorted(group_artifacts[index] & group_artifacts[index + 1]) - score = 0.0 - evidence: list[str] = [] - - if overlap: - score += 0.7 - evidence.extend(overlap[:4]) - - left_tool = _step_label(left_group) - right_tool = _step_label(right_group) - if left_tool and left_tool in group_texts[index + 1]: - score += 0.1 - evidence.append(left_tool) - - if any(marker in group_texts[index + 1] for marker in _TRANSITION_MARKERS): - score += 0.15 - evidence.append("transition") - - if any(token in group_texts[index + 1] for token in ("source", "target", "input", "output", "result", "from", "use")): - score += 0.1 - evidence.append("field_reference") - - if _group_transition_markers(right_group): - score += 0.05 - - if score >= 0.5: - edges.append( - { - "from": left_group[0].event_id, - "to": right_group[0].event_id, - "from_tool": left_tool, - "to_tool": right_tool, - "score": min(score, 1.0), - "evidence": evidence, - } - ) - return edges - - -def _workflow_content(prefix: str, groups: list[list[MemoryEvent]], dependencies: list[dict[str, Any]]) -> str: - steps = [f"{index + 1}. {_step_label(group)}" for index, group in enumerate(groups)] - if dependencies: - edges = " -> ".join(edge["from_tool"] + "→" + edge["to_tool"] for edge in dependencies if edge.get("from_tool") and edge.get("to_tool")) - if edges: - return f"{prefix}: " + " | ".join(steps) + f" | deps: {edges}" - return f"{prefix}: " + " | ".join(steps) - - -def _workflow_confidence(group_count: int, dependency_count: int, complex_flow: bool) -> float: - confidence = 0.62 - confidence += min(0.12, 0.04 * max(0, group_count - 2)) - confidence += min(0.12, 0.06 * dependency_count) - if complex_flow: - confidence += 0.08 - return min(confidence, 0.95) - - -def _workflow_reproduction_rate(groups: list[list[MemoryEvent]], dependencies: list[dict[str, Any]]) -> tuple[float, dict[str, float]]: - """Score whether a workflow contains enough evidence to replay safely. - - A workflow is reproducible only when its ordered tool steps, adjacent data - dependencies, and terminal tool results are all present. This is an - evidence score, not a fabricated constant, and is intentionally surfaced - to downstream storage and evaluation code. - """ - if not groups: - return 0.0, {"step_coverage": 0.0, "dependency_coverage": 0.0, "result_coverage": 0.0} - - step_coverage = sum(bool(_step_label(group)) for group in groups) / len(groups) - expected_dependencies = max(0, len(groups) - 1) - dependency_coverage = 1.0 if expected_dependencies == 0 else min(1.0, len(dependencies) / expected_dependencies) - result_coverage = sum( - any(event.event_type is EventType.TOOL_RESULT and event.success is not False for event in group) - for group in groups - ) / len(groups) - rate = 0.35 * step_coverage + 0.45 * dependency_coverage + 0.20 * result_coverage - return round(min(rate, 1.0), 4), { - "step_coverage": round(step_coverage, 4), - "dependency_coverage": round(dependency_coverage, 4), - "result_coverage": round(result_coverage, 4), - } - - -def _workflow_candidate( - *, - pattern: str, - prefix: str, - events: list[MemoryEvent], - groups: list[list[MemoryEvent]], - dependencies: list[dict[str, Any]], -) -> MemoryCandidate: - step_labels = [_step_label(group) for group in groups] - source_event_ids = [event.event_id for event in events] - source_summaries = [event.content or event.tool_name or event.event_type.value for event in events] - reproduction_rate, reconstruction_evidence = _workflow_reproduction_rate(groups, dependencies) - confidence = _workflow_confidence(len(groups), len(dependencies), pattern != "tool_sequence") - signature = "__".join(_slugify(label) for label in step_labels) - boundary = (0, len(events) - 1) if events else (0, 0) - - return MemoryCandidate( - candidate_id=_stable_candidate_id(events[0].user_id, MemoryType.WORKFLOW, f"workflow.{pattern}.{signature}"), - user_id=events[0].user_id, - memory_type=MemoryType.WORKFLOW, - key=f"workflow.{pattern}.{signature}", - content=_workflow_content(prefix, groups, dependencies), - scenario=_segment_scenario(events), - confidence=confidence, - source=f"workflow_{pattern}", - source_events=source_event_ids, - source_summaries=source_summaries, - tags=["workflow", pattern] + (["dependency"] if dependencies else []), - metadata={ - "pattern": pattern, - "step_count": len(groups), - "tool_names": step_labels, - "dependencies": dependencies, - "reproduction_rate": reproduction_rate, - "reconstruction_evidence": reconstruction_evidence, - "occurrence_count": 1, - "task_ids": [events[0].task_id], - "boundary": boundary, - "group_event_counts": [len(group) for group in groups], - "workflow_signature": signature, - }, - ) - - -def _extract_candidates_from_segment(segment: list[MemoryEvent], *, mode: str) -> list[MemoryCandidate]: - groups = _tool_groups(segment) - if len(groups) < 2: - return [] - - dependencies = _dependency_edges(groups) - has_conversation_glue = any(not _is_tool_event(event) for event in segment) - distinct_tools = len({label for label in (_step_label(group) for group in groups)}) - - if mode == "tool_sequence": - candidate = _workflow_candidate( - pattern="tool_sequence", - prefix="Tool sequence", - events=segment, - groups=groups, - dependencies=dependencies, - ) - return [candidate] if candidate.metadata["reproduction_rate"] >= 0.8 else [] - - complex_flow = distinct_tools >= 2 or dependencies or has_conversation_glue or len(groups) >= 3 - if not complex_flow: - return [] - - candidate = _workflow_candidate( - pattern="multi_step", - prefix="Complex workflow", - events=segment, - groups=groups, - dependencies=dependencies, - ) - return [candidate] if candidate.metadata["reproduction_rate"] >= 0.8 else [] - +from __future__ import annotations -def _dedupe_candidates(candidates: list[MemoryCandidate]) -> list[MemoryCandidate]: - unique: dict[tuple[str, str, str], MemoryCandidate] = {} - for candidate in candidates: - signature = (candidate.user_id, candidate.memory_type.value, candidate.key) - existing = unique.get(signature) - if existing is None: - unique[signature] = candidate - continue +from core.constants import MemoryType +from core.models import MemoryCandidate, MemoryEvent - merged_events = sorted(set(existing.source_events) | set(candidate.source_events)) - merged_summaries = list(dict.fromkeys(existing.source_summaries + candidate.source_summaries)) - old_count = int(existing.metadata.get("occurrence_count", 1)) - new_count = int(candidate.metadata.get("occurrence_count", 1)) - weighted_rate = ( - float(existing.metadata.get("reproduction_rate", 0.0)) * old_count - + float(candidate.metadata.get("reproduction_rate", 0.0)) * new_count - ) / (old_count + new_count) - metadata = dict(existing.metadata) - metadata["occurrence_count"] = old_count + new_count - metadata["reproduction_rate"] = round(weighted_rate, 4) - metadata["task_ids"] = sorted( - set(existing.metadata.get("task_ids", [])) | set(candidate.metadata.get("task_ids", [])) - ) - unique[signature] = MemoryCandidate( - candidate_id=existing.candidate_id, - user_id=existing.user_id, - memory_type=existing.memory_type, - key=existing.key, - content=existing.content, - scenario=existing.scenario, - confidence=max(existing.confidence, candidate.confidence), - source=existing.source, - source_events=merged_events, - source_summaries=merged_summaries, - tags=list(dict.fromkeys(existing.tags + candidate.tags)), - metadata=metadata, - created_at=existing.created_at, - ) - return list(unique.values()) +from .llm_memory_extractor import detect_boundaries_with_default_client, extract_candidates_with_default_client class WorkflowExtractor: @staticmethod def detect_workflow_boundary(events: list[MemoryEvent]) -> list[tuple[int, int]]: - return _workflow_boundaries_with_indices(events) + return detect_boundaries_with_default_client(events) @staticmethod def extract_tool_sequence(events: list[MemoryEvent]) -> list[MemoryCandidate]: - candidates: list[MemoryCandidate] = [] - for _, group_events in _group_events(events): - for start, end in _workflow_boundaries_with_indices(group_events): - segment = group_events[start : end + 1] - candidates.extend(_extract_candidates_from_segment(segment, mode="tool_sequence")) - return _dedupe_candidates(candidates) + return extract_candidates_with_default_client( + events, + mode="workflow_tool_sequence", + memory_types={MemoryType.WORKFLOW}, + ) @staticmethod def extract_multi_step_workflow(session_events: list[MemoryEvent]) -> list[MemoryCandidate]: - candidates: list[MemoryCandidate] = [] - for _, group_events in _group_events(session_events): - for start, end in _workflow_boundaries_with_indices(group_events): - segment = group_events[start : end + 1] - candidates.extend(_extract_candidates_from_segment(segment, mode="multi_step")) - return _dedupe_candidates(candidates) + return extract_candidates_with_default_client( + session_events, + mode="workflow_multi_step", + memory_types={MemoryType.WORKFLOW}, + ) diff --git a/memory/admission.py b/memory/admission.py new file mode 100644 index 0000000..deea55a --- /dev/null +++ b/memory/admission.py @@ -0,0 +1,578 @@ +"""Candidate-to-record admission with semantic conflict decisions. + +This module is the write boundary between B-side extractors and long-term +storage. LLM output decides semantic relationships; deterministic code owns +validation, idempotency, lifecycle transitions, versioning, and SQLite +transactions. +""" + +from __future__ import annotations + +import copy +import sqlite3 +from contextlib import closing +from dataclasses import dataclass +from datetime import datetime +from pathlib import Path +from typing import Iterable + +from core.constants import MemoryStatus, MemoryType, Scene +from core.models import MemoryCandidate, MemoryRecord +from extractors.llm_memory_extractor import CandidateValidator + +from .conflict_resolver import ( + ConflictAction, + ConflictDecision, + LLMConflictResolver, +) +from .lifecycle_state import ensure_transition_allowed +from .store import init_db +from .version_manager import memory_id_for_candidate, next_version + + +_RECORD_COLUMNS = ( + "memory_id, user_id, memory_type, key, content, scenario, confidence, " + "version, status, source, created_at, updated_at" +) + + +class ConcurrentAdmissionError(RuntimeError): + """The active-memory snapshot changed before the transaction committed.""" + + +@dataclass(frozen=True) +class RepositoryAdmissionOutcome: + action: ConflictAction + record: MemoryRecord | None + previous_records: tuple[MemoryRecord, ...] + + +@dataclass(frozen=True) +class AdmissionResult: + """Observable result returned to callers, demos, and audit integrations.""" + + candidate_id: str + action: ConflictAction + status: MemoryStatus + reason: str + record: MemoryRecord | None = None + previous_records: tuple[MemoryRecord, ...] = () + decision: ConflictDecision | None = None + retry_count: int = 0 + + def to_dict(self) -> dict: + return { + "candidate_id": self.candidate_id, + "action": self.action.value, + "status": self.status.value, + "reason": self.reason, + "record": self.record.to_dict() if self.record else None, + "previous_records": [record.to_dict() for record in self.previous_records], + "decision": self.decision.to_dict() if self.decision else None, + "retry_count": self.retry_count, + } + + +class SQLiteAdmissionRepository: + """Transactional adapter over the existing Phase 0 ``memories`` table. + + No table or column is added. WAL mode and a busy timeout improve concurrent + reader/writer behaviour without changing the schema. + """ + + def __init__(self, db_path: str, *, busy_timeout_ms: int = 5000) -> None: + self.db_path = db_path + self.busy_timeout_ms = max(0, int(busy_timeout_ms)) + init_db(db_path) + self._configure_database() + + def _connect(self) -> sqlite3.Connection: + Path(self.db_path).parent.mkdir(parents=True, exist_ok=True) + connection = sqlite3.connect( + self.db_path, + timeout=self.busy_timeout_ms / 1000, + isolation_level=None, + ) + connection.row_factory = sqlite3.Row + connection.execute(f"PRAGMA busy_timeout = {self.busy_timeout_ms}") + return connection + + def _configure_database(self) -> None: + with closing(self._connect()) as connection: + connection.execute("PRAGMA journal_mode = WAL") + connection.execute("PRAGMA synchronous = NORMAL") + + def get(self, memory_id: str) -> MemoryRecord | None: + with closing(self._connect()) as connection: + row = connection.execute( + f"SELECT {_RECORD_COLUMNS} FROM memories WHERE memory_id = ?", + (memory_id,), + ).fetchone() + return _row_to_record(row) if row else None + + def list_related( + self, + candidate: MemoryCandidate, + *, + statuses: Iterable[MemoryStatus] | None = None, + ) -> list[MemoryRecord]: + """List the same user/type/key stream across scenarios.""" + + with closing(self._connect()) as connection: + return self._list_related_in_connection(connection, candidate, statuses=statuses) + + def _list_related_in_connection( + self, + connection: sqlite3.Connection, + candidate: MemoryCandidate, + *, + statuses: Iterable[MemoryStatus] | None = None, + ) -> list[MemoryRecord]: + query = ( + f"SELECT {_RECORD_COLUMNS} FROM memories " + "WHERE user_id = ? AND memory_type = ? AND key = ?" + ) + params: list[object] = [ + candidate.user_id, + candidate.memory_type.value, + candidate.key, + ] + status_values = [status.value for status in statuses or ()] + if status_values: + placeholders = ",".join("?" for _ in status_values) + query += f" AND status IN ({placeholders})" + params.extend(status_values) + query += " ORDER BY scenario, version, created_at, memory_id" + rows = connection.execute(query, params).fetchall() + return [_row_to_record(row) for row in rows] + + def list_all(self, *, user_id: str | None = None) -> list[MemoryRecord]: + query = f"SELECT {_RECORD_COLUMNS} FROM memories" + params: tuple[object, ...] = () + if user_id is not None: + query += " WHERE user_id = ?" + params = (user_id,) + query += " ORDER BY created_at, memory_id" + with closing(self._connect()) as connection: + rows = connection.execute(query, params).fetchall() + return [_row_to_record(row) for row in rows] + + def apply( + self, + candidate: MemoryCandidate, + decision: ConflictDecision, + *, + expected_active_ids: tuple[str, ...], + ) -> RepositoryAdmissionOutcome: + """Atomically verify the snapshot and apply one validated decision.""" + + connection = self._connect() + try: + connection.execute("BEGIN IMMEDIATE") + active_records = self._list_related_in_connection( + connection, + candidate, + statuses=(MemoryStatus.ACTIVE,), + ) + current_ids = tuple(sorted(record.memory_id for record in active_records)) + if current_ids != tuple(sorted(expected_active_ids)): + raise ConcurrentAdmissionError( + "active memory set changed during semantic decision" + ) + previous_records = tuple(copy.deepcopy(active_records)) + active_ids = {record.memory_id for record in active_records} + target_ids = set(decision.target_memory_ids) + if decision.action in { + ConflictAction.DUPLICATE, + ConflictAction.MERGE, + ConflictAction.REPLACE, + } and (not target_ids or not target_ids.issubset(active_ids)): + raise ConcurrentAdmissionError( + "semantic decision targets are no longer active" + ) + + history = self._list_related_in_connection(connection, candidate) + base_id = memory_id_for_candidate(candidate) + base_row = connection.execute( + f"SELECT {_RECORD_COLUMNS} FROM memories WHERE memory_id = ?", + (base_id,), + ).fetchone() + base_record = _row_to_record(base_row) if base_row else None + + if decision.action is ConflictAction.DUPLICATE: + target = self._target_record(active_records, decision.target_memory_ids[0]) + if base_record and base_record.status is MemoryStatus.PENDING: + self._set_status(connection, base_record, MemoryStatus.REJECTED) + connection.commit() + return RepositoryAdmissionOutcome( + ConflictAction.DUPLICATE, + target, + previous_records, + ) + + if decision.action is ConflictAction.REJECT: + rejected = None + if base_record and base_record.status is MemoryStatus.PENDING: + rejected = self._set_status( + connection, + base_record, + MemoryStatus.REJECTED, + ) + connection.commit() + return RepositoryAdmissionOutcome( + ConflictAction.REJECT, + rejected, + previous_records, + ) + + if decision.action is ConflictAction.PENDING: + if base_record: + connection.commit() + return RepositoryAdmissionOutcome( + ConflictAction.PENDING, + base_record, + previous_records, + ) + pending = self._insert_record( + connection, + candidate, + memory_id=base_id, + content=candidate.content, + version=next_version(history, candidate), + status=MemoryStatus.PENDING, + ) + connection.commit() + return RepositoryAdmissionOutcome( + ConflictAction.PENDING, + pending, + previous_records, + ) + + final_content = decision.final_content.strip() + final_id = memory_id_for_candidate(candidate, content=final_content) + final_row = connection.execute( + f"SELECT {_RECORD_COLUMNS} FROM memories WHERE memory_id = ?", + (final_id,), + ).fetchone() + final_record = _row_to_record(final_row) if final_row else None + + if final_record and final_record.status is not MemoryStatus.PENDING: + connection.commit() + return RepositoryAdmissionOutcome( + ConflictAction.DUPLICATE, + final_record, + previous_records, + ) + + if decision.action in {ConflictAction.MERGE, ConflictAction.REPLACE}: + for record in active_records: + if record.memory_id in target_ids: + self._set_status(connection, record, MemoryStatus.SUPERSEDED) + + if base_record and base_record.status is MemoryStatus.PENDING and base_id != final_id: + self._set_status(connection, base_record, MemoryStatus.REJECTED) + + if final_record and final_record.status is MemoryStatus.PENDING: + admitted = self._activate_pending( + connection, + final_record, + candidate, + content=final_content, + ) + else: + admitted = self._insert_record( + connection, + candidate, + memory_id=final_id, + content=final_content, + version=next_version(history, candidate), + status=MemoryStatus.ACTIVE, + ) + + if decision.action in {ConflictAction.MERGE, ConflictAction.REPLACE}: + admitted.supersedes = decision.target_memory_ids[0] + connection.commit() + return RepositoryAdmissionOutcome( + decision.action, + admitted, + previous_records, + ) + except Exception: + connection.rollback() + raise + finally: + connection.close() + + def transition( + self, + memory_id: str, + target_status: MemoryStatus, + ) -> MemoryRecord: + """Apply one lifecycle transition under a write transaction.""" + + connection = self._connect() + try: + connection.execute("BEGIN IMMEDIATE") + row = connection.execute( + f"SELECT {_RECORD_COLUMNS} FROM memories WHERE memory_id = ?", + (memory_id,), + ).fetchone() + if not row: + raise KeyError(f"memory not found: {memory_id}") + current = _row_to_record(row) + ensure_transition_allowed(current.status, target_status) + updated = self._set_status(connection, current, target_status) + connection.commit() + return updated + except Exception: + connection.rollback() + raise + finally: + connection.close() + + @staticmethod + def _target_record( + records: list[MemoryRecord], + memory_id: str, + ) -> MemoryRecord: + for record in records: + if record.memory_id == memory_id: + return record + raise ConcurrentAdmissionError(f"target memory is no longer active: {memory_id}") + + @staticmethod + def _set_status( + connection: sqlite3.Connection, + record: MemoryRecord, + target_status: MemoryStatus, + ) -> MemoryRecord: + ensure_transition_allowed(record.status, target_status) + if record.status == target_status: + return record + now = datetime.now() + deleted_at = now.isoformat() if target_status is MemoryStatus.DELETED else None + connection.execute( + "UPDATE memories SET status = ?, updated_at = ?, deleted_at = ? WHERE memory_id = ?", + (target_status.value, now.isoformat(), deleted_at, record.memory_id), + ) + record.status = target_status + record.updated_at = now + return record + + @staticmethod + def _insert_record( + connection: sqlite3.Connection, + candidate: MemoryCandidate, + *, + memory_id: str, + content: str, + version: int, + status: MemoryStatus, + ) -> MemoryRecord: + now = datetime.now() + connection.execute( + """ + INSERT INTO memories + (memory_id, user_id, memory_type, key, content, scenario, + confidence, version, status, source, created_at, updated_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + """, + ( + memory_id, + candidate.user_id, + candidate.memory_type.value, + candidate.key, + content, + candidate.scenario.value, + candidate.confidence, + version, + status.value, + candidate.source, + now.isoformat(), + now.isoformat(), + ), + ) + return MemoryRecord( + memory_id=memory_id, + user_id=candidate.user_id, + memory_type=candidate.memory_type, + key=candidate.key, + content=content, + scenario=candidate.scenario, + confidence=candidate.confidence, + version=version, + status=status, + source=candidate.source, + source_events=list(candidate.source_events), + source_summaries=list(candidate.source_summaries), + tags=list(candidate.tags), + metadata=dict(candidate.metadata), + created_at=now, + updated_at=now, + ) + + @staticmethod + def _activate_pending( + connection: sqlite3.Connection, + record: MemoryRecord, + candidate: MemoryCandidate, + *, + content: str, + ) -> MemoryRecord: + ensure_transition_allowed(record.status, MemoryStatus.ACTIVE) + now = datetime.now() + connection.execute( + """ + UPDATE memories + SET content = ?, confidence = ?, source = ?, status = ?, updated_at = ? + WHERE memory_id = ? + """, + ( + content, + candidate.confidence, + candidate.source, + MemoryStatus.ACTIVE.value, + now.isoformat(), + record.memory_id, + ), + ) + record.content = content + record.confidence = candidate.confidence + record.source = candidate.source + record.status = MemoryStatus.ACTIVE + record.updated_at = now + record.source_events = list(candidate.source_events) + record.source_summaries = list(candidate.source_summaries) + record.tags = list(candidate.tags) + record.metadata = dict(candidate.metadata) + return record + + +class MemoryAdmissionService: + """Public B-side service that promotes candidates into formal records.""" + + def __init__( + self, + repository: SQLiteAdmissionRepository, + conflict_resolver: LLMConflictResolver, + *, + validator: CandidateValidator | None = None, + max_retries: int = 3, + ) -> None: + self.repository = repository + self.conflict_resolver = conflict_resolver + self.validator = validator or CandidateValidator() + self.max_retries = max(1, int(max_retries)) + + def admit(self, candidate: MemoryCandidate) -> AdmissionResult: + checked = self.validator.validate(candidate) + if checked is None: + return AdmissionResult( + candidate_id=candidate.candidate_id, + action=ConflictAction.REJECT, + status=MemoryStatus.REJECTED, + reason="candidate_failed_security_or_completeness_validation", + ) + + existing_id = memory_id_for_candidate(checked) + existing = self.repository.get(existing_id) + if existing and existing.status is not MemoryStatus.PENDING: + return AdmissionResult( + candidate_id=checked.candidate_id, + action=ConflictAction.DUPLICATE, + status=existing.status, + reason="idempotent_candidate_retry", + record=existing, + ) + + for retry_count in range(self.max_retries): + retry_existing = self.repository.get(memory_id_for_candidate(checked)) + if retry_existing and retry_existing.status is not MemoryStatus.PENDING: + return AdmissionResult( + candidate_id=checked.candidate_id, + action=ConflictAction.DUPLICATE, + status=retry_existing.status, + reason="idempotent_candidate_retry", + record=retry_existing, + retry_count=retry_count, + ) + active_records = self.repository.list_related( + checked, + statuses=(MemoryStatus.ACTIVE,), + ) + try: + decision = self.conflict_resolver.resolve(checked, active_records) + except Exception as exc: + return AdmissionResult( + candidate_id=checked.candidate_id, + action=ConflictAction.PENDING, + status=MemoryStatus.PENDING, + reason=f"llm_unavailable:{type(exc).__name__}", + previous_records=tuple(active_records), + retry_count=retry_count, + ) + + try: + outcome = self.repository.apply( + checked, + decision, + expected_active_ids=tuple( + record.memory_id for record in active_records + ), + ) + except ConcurrentAdmissionError: + if retry_count + 1 == self.max_retries: + raise + continue + + status = ( + outcome.record.status + if outcome.record is not None + else ( + MemoryStatus.REJECTED + if outcome.action is ConflictAction.REJECT + else MemoryStatus.PENDING + ) + ) + return AdmissionResult( + candidate_id=checked.candidate_id, + action=outcome.action, + status=status, + reason=decision.reason, + record=outcome.record, + previous_records=outcome.previous_records, + decision=decision, + retry_count=retry_count, + ) + + raise ConcurrentAdmissionError("memory admission retries exhausted") + + def admit_many( + self, + candidates: Iterable[MemoryCandidate], + ) -> list[AdmissionResult]: + return [self.admit(candidate) for candidate in candidates] + + def transition( + self, + memory_id: str, + target_status: MemoryStatus, + ) -> MemoryRecord: + return self.repository.transition(memory_id, target_status) + + +def _row_to_record(row: sqlite3.Row) -> MemoryRecord: + return MemoryRecord( + memory_id=row["memory_id"], + user_id=row["user_id"], + memory_type=MemoryType(row["memory_type"]), + key=row["key"], + content=row["content"], + scenario=Scene(row["scenario"]), + confidence=row["confidence"], + version=row["version"], + status=MemoryStatus(row["status"]), + source=row["source"] or "extracted", + created_at=datetime.fromisoformat(row["created_at"]), + updated_at=datetime.fromisoformat(row["updated_at"]), + ) diff --git a/memory/conflict_resolver.py b/memory/conflict_resolver.py index e69de29..9e1a439 100644 --- a/memory/conflict_resolver.py +++ b/memory/conflict_resolver.py @@ -0,0 +1,250 @@ +"""LLM-backed semantic conflict decisions for memory admission. + +The resolver never mutates storage. It turns a candidate plus the currently +active memories into a validated, structured decision. Database state +transitions remain deterministic and are handled by ``memory.admission``. +""" + +from __future__ import annotations + +import json +import re +from dataclasses import dataclass +from enum import Enum +from typing import Any, Protocol + +from core.models import MemoryCandidate, MemoryRecord + + +class AdmissionLLMClient(Protocol): + """Minimal JSON LLM boundary shared by production adapters and test fakes.""" + + def complete_json(self, prompt: str, schema: dict[str, Any]) -> Any: + ... + + +class ConflictAction(str, Enum): + """Supported semantic decisions for a candidate.""" + + CREATE = "create" + DUPLICATE = "duplicate" + MERGE = "merge" + REPLACE = "replace" + COEXIST = "coexist" + PENDING = "pending" + REJECT = "reject" + + +@dataclass(frozen=True) +class ConflictDecision: + """Validated LLM decision consumed by the deterministic admission layer.""" + + action: ConflictAction + reason: str + target_memory_ids: tuple[str, ...] = () + final_content: str = "" + requires_human_review: bool = False + decision_confidence: float | None = None + raw_response: Any = None + + def to_dict(self) -> dict[str, Any]: + return { + "action": self.action.value, + "reason": self.reason, + "target_memory_ids": list(self.target_memory_ids), + "final_content": self.final_content, + "requires_human_review": self.requires_human_review, + "decision_confidence": self.decision_confidence, + } + + +CONFLICT_DECISION_SCHEMA: dict[str, Any] = { + "type": "object", + "required": [ + "action", + "reason", + "target_memory_ids", + "final_content", + "requires_human_review", + ], + "properties": { + "action": { + "type": "string", + "enum": [action.value for action in ConflictAction], + }, + "reason": {"type": "string"}, + "target_memory_ids": { + "type": "array", + "items": {"type": "string"}, + }, + "final_content": {"type": "string"}, + "requires_human_review": {"type": "boolean"}, + "decision_confidence": { + "type": ["number", "null"], + "minimum": 0, + "maximum": 1, + }, + }, +} + + +_SECRET_KEY_RE = re.compile( + r"(api[_-]?key|password|passwd|secret|token|authorization|cookie|private[_-]?key)", + re.IGNORECASE, +) +_SECRET_VALUE_RE = re.compile( + r"(?i)\b(?:sk-[a-z0-9_-]+|Bearer\s+[a-z0-9._~+/-]+|" + r"[a-z0-9_-]{24,}\.[a-z0-9._-]+)\b" +) +_EMAIL_RE = re.compile(r"\b[\w.+-]+@[\w.-]+\.[a-zA-Z]{2,}\b") +_PHONE_RE = re.compile(r"(? Any: + if isinstance(value, dict): + sanitized: dict[str, Any] = {} + for key, item in value.items(): + sanitized[str(key)] = ( + "[REDACTED_SECRET]" + if _SECRET_KEY_RE.search(str(key)) + else _sanitize_for_llm(item) + ) + return sanitized + if isinstance(value, set): + return [_sanitize_for_llm(item) for item in sorted(value, key=str)] + if isinstance(value, (list, tuple)): + return [_sanitize_for_llm(item) for item in value] + if isinstance(value, str): + sanitized = _SECRET_VALUE_RE.sub("[REDACTED_SECRET]", value) + sanitized = _EMAIL_RE.sub("[REDACTED_EMAIL]", sanitized) + return _PHONE_RE.sub("[REDACTED_PHONE]", sanitized) + return value + + +def build_conflict_decision_prompt( + candidate: MemoryCandidate, + active_memories: list[MemoryRecord], +) -> str: + """Build a provider-neutral prompt containing no semantic shortcut rules.""" + + payload = { + "mode": "memory_admission_conflict_resolution", + "candidate": _sanitize_for_llm(candidate.to_dict()), + "active_memories": [ + _sanitize_for_llm(memory.to_dict()) for memory in active_memories + ], + "decision_contract": { + "create": "No equivalent or conflicting active memory exists.", + "duplicate": "The candidate expresses the same durable fact; do not write a duplicate.", + "merge": "The candidate and selected memories are complementary; return canonical final_content.", + "replace": "The candidate contradicts or updates selected memories; return the new canonical content.", + "coexist": "Both facts should remain active because their scopes or scenarios differ.", + "pending": "Evidence is insufficient or the decision needs human review.", + "reject": "The candidate must not become a stored memory.", + }, + "requirements": [ + "Use semantic meaning and evidence, not text length, keyword lists, or numeric admission thresholds.", + "Only reference memory IDs included in active_memories.", + "For duplicate, merge, or replace, target_memory_ids must not be empty.", + "For create, merge, replace, or coexist, final_content must be the canonical content to store.", + "Choose pending when the evidence does not support a safe deterministic mutation.", + "Return JSON only and follow the supplied schema.", + ], + } + return json.dumps(payload, ensure_ascii=False, sort_keys=True) + + +class LLMConflictResolver: + """Ask an injected LLM for a semantic decision and validate its contract.""" + + def __init__(self, llm_client: AdmissionLLMClient) -> None: + self.llm_client = llm_client + + def resolve( + self, + candidate: MemoryCandidate, + active_memories: list[MemoryRecord], + ) -> ConflictDecision: + prompt = build_conflict_decision_prompt(candidate, active_memories) + raw = self.llm_client.complete_json(prompt, CONFLICT_DECISION_SCHEMA) + return self.parse_decision(raw, candidate, active_memories) + + @staticmethod + def parse_decision( + raw: Any, + candidate: MemoryCandidate, + active_memories: list[MemoryRecord], + ) -> ConflictDecision: + if not isinstance(raw, dict): + return _pending_decision("llm_response_is_not_an_object", raw) + + try: + action = ConflictAction(str(raw.get("action", "")).strip().lower()) + except ValueError: + return _pending_decision("llm_action_is_invalid", raw) + + reason = str(raw.get("reason") or "").strip() + final_content = str(raw.get("final_content") or "").strip() + target_value = raw.get("target_memory_ids") + if not isinstance(target_value, list): + return _pending_decision("target_memory_ids_must_be_an_array", raw) + target_ids = tuple(dict.fromkeys(str(item) for item in target_value if str(item))) + + known_ids = {memory.memory_id for memory in active_memories} + if any(memory_id not in known_ids for memory_id in target_ids): + return _pending_decision("llm_referenced_unknown_memory", raw) + + if action in { + ConflictAction.DUPLICATE, + ConflictAction.MERGE, + ConflictAction.REPLACE, + } and not target_ids: + return _pending_decision("semantic_mutation_requires_a_target", raw) + + if action in { + ConflictAction.CREATE, + ConflictAction.MERGE, + ConflictAction.REPLACE, + ConflictAction.COEXIST, + } and not final_content: + return _pending_decision("storing_action_requires_final_content", raw) + + same_scene_active = any( + memory.scenario == candidate.scenario for memory in active_memories + ) + if action is ConflictAction.CREATE and same_scene_active: + return _pending_decision("create_would_bypass_existing_same_scope_memory", raw) + + decision_confidence: float | None = None + if raw.get("decision_confidence") is not None: + try: + value = float(raw["decision_confidence"]) + except (TypeError, ValueError): + return _pending_decision("decision_confidence_is_invalid", raw) + decision_confidence = max(0.0, min(value, 1.0)) + + requires_human_review = bool(raw.get("requires_human_review", False)) + if requires_human_review and action not in { + ConflictAction.PENDING, + ConflictAction.REJECT, + }: + return _pending_decision("llm_requested_human_review", raw) + + return ConflictDecision( + action=action, + reason=reason or "llm_decision", + target_memory_ids=target_ids, + final_content=final_content, + requires_human_review=requires_human_review, + decision_confidence=decision_confidence, + raw_response=raw, + ) + + +def _pending_decision(reason: str, raw: Any) -> ConflictDecision: + return ConflictDecision( + action=ConflictAction.PENDING, + reason=reason, + requires_human_review=True, + raw_response=raw, + ) diff --git a/memory/lifecycle_state.py b/memory/lifecycle_state.py new file mode 100644 index 0000000..8e2eaa1 --- /dev/null +++ b/memory/lifecycle_state.py @@ -0,0 +1,51 @@ +"""Deterministic lifecycle state validation shared by admission and D-side flows.""" + +from __future__ import annotations + +from core.constants import MemoryStatus + + +class InvalidMemoryTransition(ValueError): + """Raised when a memory lifecycle transition is not allowed.""" + + +_ALLOWED_TRANSITIONS: dict[MemoryStatus, frozenset[MemoryStatus]] = { + MemoryStatus.PENDING: frozenset( + {MemoryStatus.ACTIVE, MemoryStatus.REJECTED, MemoryStatus.DELETED} + ), + MemoryStatus.ACTIVE: frozenset( + { + MemoryStatus.SUPERSEDED, + MemoryStatus.ARCHIVED, + MemoryStatus.EXPIRED, + MemoryStatus.DELETED, + } + ), + MemoryStatus.SUPERSEDED: frozenset( + {MemoryStatus.ARCHIVED, MemoryStatus.DELETED} + ), + MemoryStatus.ARCHIVED: frozenset( + {MemoryStatus.ACTIVE, MemoryStatus.DELETED} + ), + MemoryStatus.EXPIRED: frozenset( + {MemoryStatus.ARCHIVED, MemoryStatus.DELETED} + ), + MemoryStatus.REJECTED: frozenset({MemoryStatus.DELETED}), + MemoryStatus.DELETED: frozenset(), +} + + +def can_transition(current: MemoryStatus, target: MemoryStatus) -> bool: + """Return whether a transition is valid; same-state writes are idempotent.""" + + return current == target or target in _ALLOWED_TRANSITIONS[current] + + +def ensure_transition_allowed( + current: MemoryStatus, + target: MemoryStatus, +) -> None: + if not can_transition(current, target): + raise InvalidMemoryTransition( + f"memory status cannot transition from {current.value} to {target.value}" + ) diff --git a/memory/version_manager.py b/memory/version_manager.py index e69de29..8aed975 100644 --- a/memory/version_manager.py +++ b/memory/version_manager.py @@ -0,0 +1,46 @@ +"""Deterministic identifiers and version calculation for admitted memories.""" + +from __future__ import annotations + +import hashlib +from collections.abc import Iterable + +from core.models import MemoryCandidate, MemoryRecord + + +def memory_id_for_candidate( + candidate: MemoryCandidate, + *, + content: str | None = None, +) -> str: + """Return a stable record ID for idempotent retries of the same fact.""" + + canonical_content = candidate.content if content is None else content + payload = "\x1f".join( + ( + candidate.user_id, + candidate.candidate_id, + candidate.memory_type.value, + candidate.key, + candidate.scenario.value, + canonical_content.strip(), + ) + ).encode("utf-8") + return f"mem_{hashlib.sha256(payload).hexdigest()[:20]}" + + +def next_version( + records: Iterable[MemoryRecord], + candidate: MemoryCandidate, +) -> int: + """Calculate the next version inside one user/type/key/scenario stream.""" + + versions = [ + record.version + for record in records + if record.user_id == candidate.user_id + and record.memory_type == candidate.memory_type + and record.key == candidate.key + and record.scenario == candidate.scenario + ] + return max(versions, default=0) + 1 diff --git a/tests/test_conflict_resolver.py b/tests/test_conflict_resolver.py index ef981e9..2e7e5b2 100644 --- a/tests/test_conflict_resolver.py +++ b/tests/test_conflict_resolver.py @@ -1,2 +1,162 @@ -"""Conflict resolver tests placeholder.""" +"""Tests for the LLM conflict-decision contract.""" +from __future__ import annotations + +import json +from datetime import datetime + +from core.constants import MemoryStatus, MemoryType, Scene +from core.models import MemoryCandidate, MemoryRecord +from memory.conflict_resolver import ( + ConflictAction, + LLMConflictResolver, + build_conflict_decision_prompt, +) + + +def candidate(**overrides) -> MemoryCandidate: + values = { + "candidate_id": "cand_pdf", + "user_id": "user-1", + "memory_type": MemoryType.PREFERENCE, + "key": "preference.export.format", + "content": "用户偏好使用 PDF 导出", + "scenario": Scene.OFFICE, + "confidence": 0.9, + "source": "llm_extracted", + "metadata": {}, + } + values.update(overrides) + return MemoryCandidate(**values) + + +def record(**overrides) -> MemoryRecord: + values = { + "memory_id": "mem_pdf", + "user_id": "user-1", + "memory_type": MemoryType.PREFERENCE, + "key": "preference.export.format", + "content": "用户偏好使用 PDF 导出", + "scenario": Scene.OFFICE, + "confidence": 0.9, + "version": 1, + "status": MemoryStatus.ACTIVE, + "source": "llm_extracted", + "created_at": datetime.now(), + "updated_at": datetime.now(), + } + values.update(overrides) + return MemoryRecord(**values) + + +def response(action: str, **overrides) -> dict: + value = { + "action": action, + "reason": "semantic decision", + "target_memory_ids": [], + "final_content": "用户偏好使用 PDF 导出", + "requires_human_review": False, + "decision_confidence": 0.95, + } + value.update(overrides) + return value + + +class FakeClient: + def __init__(self, result): + self.result = result + self.prompts: list[str] = [] + + def complete_json(self, prompt, schema): + self.prompts.append(prompt) + assert schema["properties"]["action"]["enum"] + return self.result + + +def test_prompt_redacts_secret_metadata(): + prompt = build_conflict_decision_prompt( + candidate(metadata={"api_key": "sk-do-not-send"}), + [], + ) + assert "sk-do-not-send" not in prompt + assert "[REDACTED_SECRET]" in prompt + assert json.loads(prompt)["mode"] == "memory_admission_conflict_resolution" + + +def test_prompt_redacts_contact_data(): + prompt = build_conflict_decision_prompt( + candidate(content="联系 user@example.com 或 13812345678"), + [], + ) + assert "user@example.com" not in prompt + assert "13812345678" not in prompt + assert "[REDACTED_EMAIL]" in prompt + assert "[REDACTED_PHONE]" in prompt + + +def test_resolver_returns_valid_create_decision(): + client = FakeClient(response("create")) + decision = LLMConflictResolver(client).resolve(candidate(), []) + assert decision.action is ConflictAction.CREATE + assert decision.final_content == "用户偏好使用 PDF 导出" + assert len(client.prompts) == 1 + + +def test_invalid_action_fails_closed_to_pending(): + decision = LLMConflictResolver.parse_decision( + response("overwrite_everything"), + candidate(), + [], + ) + assert decision.action is ConflictAction.PENDING + assert decision.requires_human_review is True + + +def test_unknown_target_fails_closed_to_pending(): + decision = LLMConflictResolver.parse_decision( + response("replace", target_memory_ids=["does-not-exist"]), + candidate(), + [record()], + ) + assert decision.action is ConflictAction.PENDING + assert decision.reason == "llm_referenced_unknown_memory" + + +def test_mutating_decision_requires_target(): + decision = LLMConflictResolver.parse_decision( + response("merge", target_memory_ids=[]), + candidate(), + [record()], + ) + assert decision.action is ConflictAction.PENDING + assert decision.reason == "semantic_mutation_requires_a_target" + + +def test_create_cannot_bypass_same_scene_active_memory(): + decision = LLMConflictResolver.parse_decision( + response("create"), + candidate(), + [record()], + ) + assert decision.action is ConflictAction.PENDING + assert decision.reason == "create_would_bypass_existing_same_scope_memory" + + +def test_different_scene_can_use_coexist(): + existing = record(scenario=Scene.CODING) + decision = LLMConflictResolver.parse_decision( + response("coexist"), + candidate(), + [existing], + ) + assert decision.action is ConflictAction.COEXIST + + +def test_human_review_request_cannot_mutate_storage_directly(): + decision = LLMConflictResolver.parse_decision( + response("create", requires_human_review=True), + candidate(), + [], + ) + assert decision.action is ConflictAction.PENDING + assert decision.reason == "llm_requested_human_review" diff --git a/tests/test_environment_extractor.py b/tests/test_environment_extractor.py index 720bb2a..e82043f 100644 --- a/tests/test_environment_extractor.py +++ b/tests/test_environment_extractor.py @@ -1,99 +1,72 @@ from __future__ import annotations +import json +from typing import Any + +import pytest + from core.constants import MemoryType from extractors.environment_extractor import EnvironmentExtractor +from extractors.llm_memory_extractor import clear_default_llm_client, set_default_llm_client -def test_extract_environment_from_structured_tool_output() -> None: - output = { - "user_id": "alice", - "paths": { - "Downloads": r"C:\Users\alice\Downloads", - "Documents": r"C:\Users\alice\Documents", - }, - "locale": "zh_CN.UTF-8", - "installed_software": [ - {"name": "LibreOffice", "version": "24.2"}, - {"display_name": "Firefox", "release": "126"}, - ], - "system": {"os_name": "Kylin Desktop", "version": "V11"}, - "api_key": "must-not-be-extracted", - } +class EnvironmentFakeLLMClient: + def __init__(self) -> None: + self.calls: list[dict[str, Any]] = [] - candidates = EnvironmentExtractor.extract_from_tool_output(output) - by_key = {candidate.key: candidate for candidate in candidates} - - assert set(by_key) >= { - "environment.path.downloads", - "environment.path.documents", - "environment.locale.language", - "environment.locale.region", - "environment.software.libreoffice", - "environment.software.firefox", - "environment.system.version", - } - assert all(candidate.user_id == "alice" for candidate in candidates) - assert all(candidate.memory_type is MemoryType.ENVIRONMENT for candidate in candidates) - assert by_key["environment.path.downloads"].metadata["path"] == r"~\Downloads" - assert by_key["environment.path.documents"].metadata["path"] == r"~\Documents" - assert by_key["environment.locale.language"].metadata["language"] == "zh" - assert by_key["environment.locale.region"].metadata["region"] == "CN" - assert by_key["environment.software.libreoffice"].metadata["version"] == "24.2" - assert by_key["environment.system.version"].metadata["version"] == "Kylin Desktop V11" - assert "api_key" not in str([candidate.to_dict() for candidate in candidates]) - - -def test_environment_extraction_supports_common_output_shapes_and_is_deterministic() -> None: - output = { - "downloads": "/home/bob/Downloads", - "documents": "/home/bob/Documents", - "language": "en-US", - "region": "us", - "applications": { - "VS Code": "1.90", - "git": {"version": "2.45"}, - }, - "os_version": "Ubuntu 24.04", - } - - first = EnvironmentExtractor.extract_from_tool_output(output) - second = EnvironmentExtractor.extract_from_tool_output(output) - by_key = {candidate.key: candidate for candidate in first} - - assert by_key["environment.path.downloads"].metadata["path"] == r"~\Downloads" - assert by_key["environment.path.documents"].metadata["path"] == r"~\Documents" - assert by_key["environment.locale.language"].metadata["language"] == "en" - assert by_key["environment.locale.region"].metadata["region"] == "US" - assert by_key["environment.software.vs_code"].metadata["version"] == "1.90" - assert by_key["environment.software.git"].metadata["version"] == "2.45" - assert by_key["environment.system.version"].metadata["version"] == "Ubuntu 24.04" - assert [(candidate.key, candidate.candidate_id) for candidate in first] == [ - (candidate.key, candidate.candidate_id) for candidate in second - ] - - -def test_environment_extractor_rejects_non_mapping_output() -> None: + def complete_json(self, prompt: str, schema: dict[str, Any]) -> dict[str, Any]: + request = json.loads(prompt) + self.calls.append({"request": request, "prompt": prompt, "schema": schema}) + return { + "candidates": [ + { + "is_memory_worthy": True, + "is_long_term": True, + "memory_type": "environment", + "category": "path", + "value": "documents", + "scope": "system", + "content": "常用文档目录是 /root/docs", + "confidence": 0.9, + "evidence": "documents=/root/docs", + "reason": "LLM 从系统上下文中识别出可复用目录配置。", + "sensitivity": "none", + } + ] + } + + +@pytest.fixture() +def env_client() -> EnvironmentFakeLLMClient: + client = EnvironmentFakeLLMClient() + set_default_llm_client(client) try: - EnvironmentExtractor.extract_from_tool_output([]) # type: ignore[arg-type] - except TypeError as exc: - assert "output must be a dict" in str(exc) - else: # pragma: no cover - explicit failure branch - raise AssertionError("non-mapping output must be rejected") - - -def test_environment_preserves_absolute_docs_path_and_filters_nested_secrets() -> None: - output = { - "docs": "/root/docs", - "nested": { - "token": "must-not-leak", - "paths": {"download_dir": "/root/downloads"}, - }, - } + yield client + finally: + clear_default_llm_client() + + +def test_environment_extractor_wraps_tool_output_for_llm(env_client: EnvironmentFakeLLMClient) -> None: + output = {"user_id": "u-env", "documents": "/root/docs", "api_key": "sk-secret-1234567890"} candidates = EnvironmentExtractor.extract_from_tool_output(output) - by_key = {candidate.key: candidate for candidate in candidates} - rendered = str([candidate.to_dict() for candidate in candidates]) - assert by_key["environment.path.documents"].metadata["path"] == r"\root\docs" - assert by_key["environment.path.downloads"].metadata["path"] == r"\root\downloads" - assert "must-not-leak" not in rendered + assert candidates[0].memory_type is MemoryType.ENVIRONMENT + assert candidates[0].key == "environment.path.documents" + request = env_client.calls[0]["request"] + assert request["mode"] == "environment_from_tool_output" + assert request["events"][0]["output"]["documents"] == "/root/docs" + assert request["events"][0]["output"]["api_key"] == "[REDACTED_SECRET]" + assert "sk-secret-1234567890" not in env_client.calls[0]["prompt"] + + +def test_environment_extractor_has_no_rule_fallback_without_client() -> None: + clear_default_llm_client() + + assert EnvironmentExtractor.extract_from_tool_output({"documents": "/root/docs"}) == [] + + +def test_environment_extractor_rejects_non_dict_and_empty_input(env_client: EnvironmentFakeLLMClient) -> None: + assert EnvironmentExtractor.extract_from_tool_output({}) == [] + assert EnvironmentExtractor.extract_from_tool_output([]) == [] # type: ignore[arg-type] + assert env_client.calls == [] diff --git a/tests/test_extractors.py b/tests/test_extractors.py index d5a2d18..5a901ba 100644 --- a/tests/test_extractors.py +++ b/tests/test_extractors.py @@ -1,785 +1,152 @@ from __future__ import annotations -import sqlite3 +import json from datetime import datetime, timezone -from pathlib import Path +from typing import Any + +import pytest from core.constants import EventType, MemoryType, Scene -from core.models import MemoryCandidate, MemoryEvent, MemoryRecord -from extractors import knowledge_extractor as knowledge_mod -from extractors import preference_extractor as preference_mod +from core.models import MemoryEvent from extractors.knowledge_extractor import KnowledgeExtractor +from extractors.llm_memory_extractor import clear_default_llm_client, set_default_llm_client from extractors.preference_extractor import PreferenceExtractor -from memory.store import init_db, save_memory + + +class RoutingFakeLLMClient: + def __init__(self) -> None: + self.calls: list[dict[str, Any]] = [] + + def complete_json(self, prompt: str, schema: dict[str, Any]) -> dict[str, Any]: + request = json.loads(prompt) + self.calls.append({"request": request, "schema": schema}) + mode = request["mode"] + if mode.startswith("preference"): + return _payload("preference", "output_format", "markdown", "用户偏好 Markdown 输出") + if mode.startswith("knowledge"): + return _payload("knowledge", "tool_case", "batch_export", "工具结果可复用为批量导出知识") + if mode.startswith("template"): + return _payload("template", "data_processing", "merge_files", "合并文件可复用模板") + return {"candidates": []} + + +def _payload(memory_type: str, category: str, value: str, content: str) -> dict[str, Any]: + return { + "candidates": [ + { + "is_memory_worthy": True, + "is_long_term": True, + "memory_type": memory_type, + "category": category, + "value": value, + "scope": "unit_test", + "content": content, + "confidence": 0.86, + "evidence": "fake LLM evidence", + "reason": "fake LLM reason", + "sensitivity": "none", + } + ] + } def _event( event_id: str, *, - user_id: str = "user-1", - session_id: str = "session-1", - task_id: str = "task-1", event_type: EventType = EventType.CONVERSATION, - scenario: Scene = Scene.GLOBAL, - source: str = "conversation", content: str | None = None, tool_name: str | None = None, - input_payload: dict | None = None, - output_payload: dict | None = None, - metadata: dict | None = None, - success: bool | None = True, - timestamp: datetime | None = None, + output: dict[str, Any] | None = None, ) -> MemoryEvent: return MemoryEvent( event_id=event_id, raw_event_id=f"raw-{event_id}", - user_id=user_id, - session_id=session_id, - task_id=task_id, + user_id="user-b", + session_id="session-b", + task_id="task-b", event_type=event_type, - scenario=scenario, - source=source, + scenario=Scene.OFFICE, + source=event_type.value, + actor="user", content=content, tool_name=tool_name, - input=input_payload or {}, - output=output_payload or {}, - metadata=metadata or {}, - success=success, - timestamp=timestamp or datetime(2026, 6, 23, tzinfo=timezone.utc), - raw_event={"event_id": event_id}, + output=output or {}, + timestamp=datetime(2026, 7, 5, 10, 0, tzinfo=timezone.utc), ) -def test_preference_explicit_from_conversation() -> None: - content = "以后都用 Markdown 输出,回答尽量简洁,先给结论后分析。" - - raw_candidates = PreferenceExtractor.extract_explicit_preference(content) - assert raw_candidates - assert any(candidate.key == "preference.output_format.markdown" for candidate in raw_candidates) - assert all(candidate.user_id == "" for candidate in raw_candidates) - - event = _event("evt-1", content=content) - candidates = PreferenceExtractor.extract_from_conversation(event) +@pytest.fixture() +def llm_client() -> RoutingFakeLLMClient: + client = RoutingFakeLLMClient() + set_default_llm_client(client) + try: + yield client + finally: + clear_default_llm_client() - keys = {candidate.key for candidate in candidates} - assert "preference.output_format.markdown" in keys - assert "preference.response_length.简洁" in keys - assert all(candidate.user_id == event.user_id for candidate in candidates) - assert all(event.event_id in candidate.source_events for candidate in candidates) - assert all(0.5 <= candidate.confidence <= 1.0 for candidate in candidates) - - -def test_preference_implicit_from_frequency() -> None: - events = [ - _event( - "evt-2", - event_type=EventType.TOOL_CALL, - source="workflow", - tool_name="bash", - input_payload={"format": "markdown"}, - ), - _event( - "evt-3", - event_type=EventType.TOOL_CALL, - source="workflow", - tool_name="bash", - input_payload={"format": "markdown"}, - ), - _event( - "evt-4", - event_type=EventType.TOOL_CALL, - source="workflow", - tool_name="bash", - input_payload={"format": "markdown"}, - ), - ] - candidates = PreferenceExtractor.extract_implicit_preference(events) +def test_preference_explicit_from_conversation_uses_llm(llm_client: RoutingFakeLLMClient) -> None: + event = _event("evt-pref", content="以后都用 Markdown 输出。") - assert any(candidate.key.startswith("preference.tool.bash") for candidate in candidates) - assert any(candidate.key.startswith("preference.parameter.format") for candidate in candidates) - assert all(candidate.memory_type is MemoryType.PREFERENCE for candidate in candidates) - assert all(candidate.user_id == "user-1" for candidate in candidates) - assert all(len(candidate.source_events) >= 3 for candidate in candidates) - assert all(0.5 <= candidate.confidence <= 1.0 for candidate in candidates) + candidates = PreferenceExtractor.extract_from_conversation(event) + assert llm_client.calls[0]["request"]["mode"] == "preference_from_conversation" + assert candidates[0].memory_type is MemoryType.PREFERENCE + assert candidates[0].key == "preference.output_format.markdown" -def test_preference_confidence_score() -> None: - explicit_event = _event( - "evt-5", - content="以后都用 PDF 输出,回答尽量简洁。", - ) - implicit_events = [ - _event("evt-6", event_type=EventType.TOOL_CALL, source="workflow", tool_name="bash", input_payload={"format": "pdf"}), - _event("evt-7", event_type=EventType.TOOL_CALL, source="workflow", tool_name="bash", input_payload={"format": "pdf"}), - _event("evt-8", event_type=EventType.TOOL_CALL, source="workflow", tool_name="bash", input_payload={"format": "pdf"}), - ] - explicit_candidate = PreferenceExtractor.extract_from_conversation(explicit_event)[0] - implicit_candidate = PreferenceExtractor.extract_implicit_preference(implicit_events)[0] +def test_preference_explicit_text_has_no_rule_fallback_without_client() -> None: + clear_default_llm_client() - assert explicit_candidate.confidence > implicit_candidate.confidence - assert 0.9 <= explicit_candidate.confidence <= 1.0 - assert 0.5 <= implicit_candidate.confidence <= 1.0 + assert PreferenceExtractor.extract_explicit_preference("以后都用 PDF 输出") == [] -def test_knowledge_from_tool_result() -> None: - complete_event = _event( - "evt-9", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="etl", - content="批量导出完成。", - input_payload={"operation": "batch export", "source": "orders.csv", "format": "xlsx"}, - output_payload={"status": "ok", "rows": 1200, "file": "orders_20260623.xlsx"}, - metadata={"mode": "batch"}, - ) - partial_event = _event( - "evt-10", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="etl", - content="批量导出处理中。", - input_payload={"operation": "batch export"}, - output_payload={}, - metadata={}, - ) +def test_preference_implicit_from_frequency_is_llm_session_extraction(llm_client: RoutingFakeLLMClient) -> None: + events = [_event("evt-pref-1", content="第 1 次使用某种格式"), _event("evt-pref-2", content="第 2 次使用某种格式")] - complete_candidates = KnowledgeExtractor.extract_from_tool_result(complete_event) - partial_candidates = KnowledgeExtractor.extract_from_tool_result(partial_event) - - assert len(complete_candidates) >= 1 - candidate = complete_candidates[0] - assert candidate.memory_type is MemoryType.KNOWLEDGE - assert candidate.user_id == complete_event.user_id - assert candidate.key.startswith("knowledge.tool_case.batch_export.") - assert "输入" in candidate.content - assert "输出" in candidate.content - assert complete_event.event_id in candidate.source_events - assert candidate.confidence >= 0.8 - assert candidate.confidence > partial_candidates[0].confidence - - faq_event = _event( - "evt-11", - content="问题:导出失败怎么办?\n解决方案:先检查文件是否被占用,再重试。\n桌面配置:1. 打开 settings;2. 关闭自动更新;3. 重启。", - ) - faq_candidates = KnowledgeExtractor.extract_from_conversation(faq_event) - faq_keys = {item.key for item in faq_candidates} - assert any(key.startswith("knowledge.faq.") for key in faq_keys) - assert any(key.startswith("knowledge.system.desktop_config") or key.startswith("knowledge.guide.") for key in faq_keys) - - -def test_knowledge_template_extraction() -> None: - events = [ - _event( - "evt-12", - event_type=EventType.TOOL_CALL, - source="workflow", - tool_name="exporter", - content="批量导出订单到 Excel", - input_payload={"operation": "batch export", "source": "orders_001.csv", "format": "xlsx"}, - output_payload={"file": "orders_001.xlsx", "status": "ok"}, - ), - _event( - "evt-13", - event_type=EventType.TOOL_CALL, - source="workflow", - tool_name="exporter", - content="批量导出订单到 Excel", - input_payload={"operation": "batch export", "source": "orders_002.csv", "format": "xlsx"}, - output_payload={"file": "orders_002.xlsx", "status": "ok"}, - ), - ] - - candidates = KnowledgeExtractor.extract_templates(events) - - assert len(candidates) == 1 - candidate = candidates[0] - assert candidate.memory_type is MemoryType.TEMPLATE - assert candidate.key.startswith("template.batch_export.") - assert set(candidate.source_events) == {"evt-12", "evt-13"} - assert candidate.confidence >= 0.6 - - -def test_knowledge_deduplication() -> None: - events = [ - _event( - "evt-14", - event_type=EventType.TOOL_CALL, - source="workflow", - tool_name="exporter", - content="批量导出订单到 Excel", - input_payload={"operation": "batch export", "source": "orders_001.csv", "format": "xlsx"}, - output_payload={"file": "orders_001.xlsx", "status": "ok"}, - ), - _event( - "evt-15", - event_type=EventType.TOOL_CALL, - source="workflow", - tool_name="exporter", - content="批量导出订单到 Excel", - input_payload={"operation": "batch export", "source": "orders_002.csv", "format": "xlsx"}, - output_payload={"file": "orders_002.xlsx", "status": "ok"}, - ), - _event( - "evt-16", - event_type=EventType.TOOL_CALL, - source="workflow", - tool_name="exporter", - content="批量导出订单到 Excel", - input_payload={"operation": "batch export", "source": "orders_003.csv", "format": "xlsx"}, - output_payload={"file": "orders_003.xlsx", "status": "ok"}, - ), - ] - - candidates = KnowledgeExtractor.extract_templates(events) - - assert len(candidates) == 1 - candidate = candidates[0] - assert candidate.key.startswith("template.batch_export.") - assert set(candidate.source_events) == {"evt-14", "evt-15", "evt-16"} - assert candidate.confidence >= 0.65 - - -def test_core_models_and_storage_round_trip(tmp_path: Path) -> None: - event = MemoryEvent( - event_id="evt-model-1", - raw_event_id="raw-model-1", - user_id="user-model", - session_id="session-model", - task_id="task-model", - event_type=EventType.CONVERSATION, - scenario=Scene.GLOBAL, - source="conversation", - content="模型测试", - metadata={"tuple": (1, 2)}, - timestamp=datetime(2026, 6, 23, tzinfo=timezone.utc), - ) + candidates = PreferenceExtractor.extract_implicit_preference(events) - event_dict = event.to_dict() - assert event_dict["event_type"] == "conversation" - assert event_dict["scenario"] == "global" - assert event_dict["timestamp"].startswith("2026-06-23T00:00:00") - - candidate = MemoryCandidate( - candidate_id="cand-model-1", - user_id="user-candidate", - memory_type=MemoryType.KNOWLEDGE, - key="knowledge.example", - content="候选内容", - scenario=Scene.GLOBAL, - source_events=["evt-a", "evt-b"], - source_summaries=["summary-1"], - tags=["tag-a", "tag-b"], - metadata={"tuple": (1, 2), "flags": [True, False]}, - created_at=datetime(2026, 6, 23, 1, tzinfo=timezone.utc), - ) + assert llm_client.calls[0]["request"]["mode"] == "preference_from_session" + assert candidates[0].source == "llm_extracted" - candidate_dict = candidate.to_dict() - assert candidate_dict["memory_type"] == "knowledge" - assert candidate_dict["user_id"] == "user-candidate" - assert "evt-a" in candidate_dict["source_events"] - assert candidate_dict["metadata"]["tuple"] == (1, 2) - - db_path = tmp_path / "memories.sqlite3" - init_db(str(db_path)) - persisted = save_memory(str(db_path), candidate) - with sqlite3.connect(db_path) as connection: - row = connection.execute( - "SELECT user_id, memory_type, key, content FROM memories" - ).fetchone() - assert row == ("user-candidate", "knowledge", "knowledge.example", "候选内容") - assert persisted.memory_id - - -def test_preference_internal_helper_branches() -> None: - assert preference_mod._normalize_value("response_style", "简洁") == ("concise", "简洁") - assert preference_mod._normalize_value("language", "English") == ("en", "English") - assert preference_mod._normalize_value("tool", "ripgrep") == ("rg", "ripgrep") - assert preference_mod._normalize_value("workflow", "先给方案")[1] == "先给方案" - assert preference_mod._normalize_value("unknown", "自定义值")[0] == "自定义值" - - assert preference_mod._canonical_structured_key("preferred_format") == "output_format" - assert preference_mod._canonical_structured_key("default_language") == "language" - assert preference_mod._canonical_structured_key("response_tone") == "response_style" - assert preference_mod._canonical_structured_key("preferred_tool") == "tool" - assert preference_mod._canonical_structured_key("response_length") == "response_length" - assert preference_mod._canonical_structured_key("workflow_mode") == "workflow" - assert preference_mod._canonical_structured_key("parameter_option") == "parameter" - assert preference_mod._canonical_structured_key("custom") == "custom" - - assert preference_mod._iter_text_fragments(None) == [] - assert preference_mod._iter_text_fragments(" abc ") == ["abc"] - assert preference_mod._iter_text_fragments(3) == ["3"] - fragments = preference_mod._iter_text_fragments( - { - "event_id": "skip", - "nested": {"x": "y"}, - "items": [1, 2], - "flags": (True, False), - } - ) - assert "y" in fragments - assert "1" in fragments and "2" in fragments - assert "True" in fragments and "False" in fragments - assert {"alpha", "beta"} == set(preference_mod._iter_text_fragments({"alpha", "beta"})) - - explicit = preference_mod._extract_from_text("") - assert explicit == [] - explicit = preference_mod._extract_from_text( - "以后都用 Markdown 输出,之后也用 Markdown 输出,偏好中文,优先使用 git,不喜欢冗长,今后先列清单,回答详细,始终分点。" - ) - assert len(explicit) >= 6 - assert any(item.key.startswith("preference.output_format") for item in explicit) - assert any(item.key.startswith("preference.language") for item in explicit) - assert any(item.key.startswith("preference.tool") for item in explicit) - assert any(item.key.startswith("preference.workflow") for item in explicit) +def test_knowledge_from_tool_result_uses_llm(llm_client: RoutingFakeLLMClient) -> None: event = _event( - "evt-pref-internal", - content="偏好测试", - input_payload={ - "preferred_format": "PDF", - "default_language": "中文", - "response_tone": "正式", - "preferred_tool": "git", - "response_length": "long", - "workflow_mode": "setup", - "parameter_option": "fast", - }, - output_payload={"preferred_format": "PDF", "default_language": "中文"}, - metadata={"preferred_format": "PDF", "response_style": "简洁"}, - ) - structured = preference_mod._structured_pref_candidates(event) - structured_keys = {item.key for item in structured} - assert "preference.output_format.pdf" in structured_keys - assert "preference.language.zh" in structured_keys - assert "preference.response_style.formal" in structured_keys - assert "preference.tool.git" in structured_keys - assert "preference.response_length.long" in structured_keys - assert "preference.workflow.setup" in structured_keys - assert "preference.parameter.fast" in structured_keys - - rebound = preference_mod._rebind_candidates(explicit[:1], event=event, source="conversation") - assert rebound[0].user_id == event.user_id - assert rebound[0].source == "conversation" - assert event.event_id in rebound[0].source_events - - deduped = preference_mod._dedupe_candidates( - [ - MemoryCandidate(candidate_id="cand-pref-a", user_id="u", memory_type=MemoryType.PREFERENCE, key="preference.test.x", content="a"), - MemoryCandidate(candidate_id="cand-pref-b", user_id="u", memory_type=MemoryType.PREFERENCE, key="preference.test.x", content="b"), - ] - ) - assert len(deduped) == 1 - - -def test_preference_implicit_internal_branches() -> None: - assert preference_mod._normalize_scalar(True) == "true" - assert preference_mod._normalize_scalar(3) == "3" - assert preference_mod._normalize_scalar(3.14) == "3.14" - assert preference_mod._normalize_scalar(" hi ") == "hi" - - params = preference_mod._collect_scalar_params( - { - "skip": None, - "nested": { - "temperature": 0.2, - "top_p": 0.9, - "flags": [True, False], - "model": "gpt-4", - }, - "tags": ("alpha", "beta"), - } - ) - params_map = dict(params) - assert params_map["nested.temperature"] == "0.2" - assert params_map["nested.top_p"] == "0.9" - assert params_map["nested.flags"] == "true, false" - assert params_map["nested.model"] == "gpt-4" - assert params_map["tags"] == "alpha, beta" - - content_event = _event("evt-sig-1", content=" 重复 操作 ") - tool_event = _event("evt-sig-2", content="", tool_name="Bash") - meta_event = _event("evt-sig-3", content="", tool_name=None, metadata={"action": "Open File"}) - fallback_event = _event("evt-sig-4", content="", tool_name=None, metadata={}) - assert preference_mod._action_signature(content_event) == "重复 操作" - assert preference_mod._action_signature(tool_event) == "tool:bash" - assert preference_mod._action_signature(meta_event).startswith("action:open file") - assert preference_mod._action_signature(fallback_event) == "conversation:conversation" - - mixed_scenario_events = [ - _event( - "evt-imp-1", - event_type=EventType.TOOL_CALL, - scenario=Scene.CODING, - source="workflow", - tool_name="bash", - input_payload={"format": "markdown", "nested": {"temperature": 0.2}}, - ), - _event( - "evt-imp-2", - event_type=EventType.TOOL_CALL, - scenario=Scene.CODING, - source="workflow", - tool_name="bash", - input_payload={"format": "markdown", "nested": {"temperature": 0.2}}, - ), - _event( - "evt-imp-3", - event_type=EventType.TOOL_CALL, - scenario=Scene.CODING, - source="workflow", - tool_name="bash", - input_payload={"format": "markdown", "nested": {"temperature": 0.2}}, - ), - _event( - "evt-imp-4", - event_type=EventType.TOOL_CALL, - scenario=Scene.GLOBAL, - source="workflow", - tool_name="git", - input_payload={"format": "markdown"}, - ), - _event( - "evt-imp-5", - event_type=EventType.TOOL_CALL, - scenario=Scene.GLOBAL, - source="workflow", - tool_name="git", - input_payload={"format": "markdown"}, - ), - _event( - "evt-imp-6", - event_type=EventType.TOOL_CALL, - scenario=Scene.GLOBAL, - source="workflow", - tool_name="git", - input_payload={"format": "markdown"}, - ), - ] - - implicit = preference_mod.PreferenceExtractor.extract_implicit_preference(mixed_scenario_events) - keys = {candidate.key for candidate in implicit} - assert any(key.startswith("preference.tool.bash") for key in keys) - assert any(key.startswith("preference.tool.git") for key in keys) - assert any(key.startswith("preference.parameter.format") for key in keys) - assert any(candidate.scenario is Scene.GLOBAL for candidate in implicit) - assert any(candidate.scenario is Scene.CODING for candidate in implicit) - assert all(candidate.user_id == "user-1" for candidate in implicit) - - tool_group = [event for event in mixed_scenario_events if event.tool_name == "bash"] - results = preference_mod._extract_implicit_for_user("user-x", tool_group) - assert results - - -def test_knowledge_internal_helper_branches() -> None: - assert knowledge_mod._text_fragments(None) == [] - assert knowledge_mod._text_fragments(" abc ") == ["abc"] - assert knowledge_mod._text_fragments(5) == ["5"] - fragments = knowledge_mod._text_fragments( - { - "event_id": "skip", - "nested": {"x": "y"}, - "items": [1, 2], - "set_items": {"alpha", "beta"}, - } - ) - assert "y" in fragments - assert "1" in fragments and "2" in fragments - assert {"alpha", "beta"}.issubset(set(fragments)) - - normalized = knowledge_mod._normalize_for_template( - r"Visit https://example.com orders_001.xlsx C:\tmp\app_2.log /var/log/app.txt 42 deadbeef 'quoted text'" - ) - assert "" in normalized - assert "" in normalized - assert "" in normalized - assert "" in normalized - assert "" in normalized - - summary = knowledge_mod._summarize_mapping( - {"a": 1, "b": [2, 3], "c": None, "d": {"skip": 1}, "e": "z"}, - max_items=3, - ) - assert summary == "a=1; b=2, 3; e=z" - - batch_event = _event( - "evt-know-intent-1", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="etl", - content="batch export", - input_payload={"operation": "batch export"}, - ) - merge_event = _event( - "evt-know-intent-2", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="merge", - content="merge files", - input_payload={"operation": "merge files"}, - ) - desktop_event = _event( - "evt-know-intent-3", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="desktop", - content="desktop config", - ) - software_event = _event( - "evt-know-intent-4", + "evt-knowledge", event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="installer", - content="software setup", - ) - generic_event = _event("evt-know-intent-5", event_type=EventType.TOOL_RESULT, source="tool", tool_name="other", content="misc") - assert knowledge_mod._extract_tool_intent(batch_event) == "batch_export" - assert knowledge_mod._extract_tool_intent(merge_event) == "merge_files" - assert knowledge_mod._extract_tool_intent(desktop_event) == "desktop_config" - assert knowledge_mod._extract_tool_intent(software_event) == "software_setup" - assert knowledge_mod._extract_tool_intent(generic_event) == "tool_case" - - assert knowledge_mod._completeness_score( - has_input=True, has_output=True, has_content=True, has_metadata=True - ) > knowledge_mod._completeness_score( - has_input=True, has_output=False, has_content=False, has_metadata=False + tool_name="exporter", + output={"status": "ok", "file": "report.pdf"}, ) - faq_candidates = knowledge_mod._extract_faq_candidates( - _event("evt-faq-1", content="Q: How to export?\nA: Use batch export.") - , "Q: How to export?\nA: Use batch export.") - assert len(faq_candidates) == 1 - - guide_candidates = knowledge_mod._extract_guide_candidates( - _event("evt-guide-1", content="步骤1: 打开软件;然后配置;最后运行。"), - "步骤1: 打开软件;然后配置;最后运行。", - ) - assert guide_candidates + candidates = KnowledgeExtractor.extract_from_tool_result(event) - system_candidates = knowledge_mod._extract_system_guide( - _event("evt-system-1", content="系统配置说明") - ) - assert system_candidates - assert system_candidates[0].key.startswith("knowledge.system.system_setup") - - -def test_knowledge_template_variants_and_dedup() -> None: - events = [ - _event( - "evt-temp-1", - event_type=EventType.TOOL_CALL, - scenario=Scene.GLOBAL, - source="workflow", - tool_name="exporter", - content="批量导出订单到 Excel", - input_payload={"operation": "batch export", "source": "orders_001.csv", "format": "xlsx"}, - output_payload={"file": "orders_001.xlsx", "status": "ok"}, - ), - _event( - "evt-temp-1b", - event_type=EventType.TOOL_CALL, - scenario=Scene.GLOBAL, - source="workflow", - tool_name="exporter", - content="批量导出订单到 Excel", - input_payload={"operation": "batch export", "source": "orders_002.csv", "format": "xlsx"}, - output_payload={"file": "orders_002.xlsx", "status": "ok"}, - ), - _event( - "evt-temp-2", - event_type=EventType.TOOL_CALL, - scenario=Scene.GLOBAL, - source="workflow", - tool_name="merger", - content="合并文件", - input_payload={"operation": "merge files", "source": "left_001.csv", "target": "right_001.csv"}, - output_payload={"file": "merged_001.csv"}, - ), - _event( - "evt-temp-2b", - event_type=EventType.TOOL_CALL, - scenario=Scene.CODING, - source="workflow", - tool_name="merger", - content="合并文件", - input_payload={"operation": "merge files", "source": "left_002.csv", "target": "right_002.csv"}, - output_payload={"file": "merged_002.csv"}, - ), - _event( - "evt-temp-3", - event_type=EventType.TOOL_CALL, - scenario=Scene.SYSTEM, - source="workflow", - tool_name="desktop", - content="桌面配置", - input_payload={"mode": "desktop config", "setting": "dark"}, - output_payload={"status": "ok"}, - ), - _event( - "evt-temp-3b", - event_type=EventType.TOOL_CALL, - scenario=Scene.SYSTEM, - source="workflow", - tool_name="desktop", - content="桌面配置", - input_payload={"mode": "desktop config", "setting": "dark"}, - output_payload={"status": "ok"}, - ), - _event( - "evt-temp-4", - event_type=EventType.TOOL_CALL, - scenario=Scene.CODING, - source="workflow", - tool_name="installer", - content="软件安装", - input_payload={"mode": "software setup", "package": "app.exe"}, - output_payload={"status": "ok"}, - ), - _event( - "evt-temp-4b", - event_type=EventType.TOOL_CALL, - scenario=Scene.CODING, - source="workflow", - tool_name="installer", - content="软件安装", - input_payload={"mode": "software setup", "package": "app.exe"}, - output_payload={"status": "ok"}, - ), - _event( - "evt-temp-5", - event_type=EventType.TOOL_CALL, - scenario=Scene.GLOBAL, - source="workflow", - tool_name="misc", - content="通用流程", - input_payload={"step": "1"}, - output_payload={"result": "done"}, - ), - _event( - "evt-temp-5b", - event_type=EventType.TOOL_CALL, - scenario=Scene.GLOBAL, - source="workflow", - tool_name="misc", - content="通用流程", - input_payload={"step": "1"}, - output_payload={"result": "done"}, - ), - ] - - candidates = knowledge_mod.KnowledgeExtractor.extract_templates(events) - keys = {candidate.key for candidate in candidates} - assert any(key.startswith("template.batch_export.") for key in keys) - assert any(key.startswith("template.merge_files.") for key in keys) - assert any(key.startswith("template.desktop_config.") for key in keys) - assert any(key.startswith("template.software_setup.") for key in keys) - assert any(key.startswith("template.generic.") for key in keys) - assert any(candidate.scenario is Scene.GLOBAL for candidate in candidates) - assert any(candidate.scenario is Scene.SYSTEM for candidate in candidates) - assert any(candidate.scenario is Scene.CODING for candidate in candidates) - - deduped = knowledge_mod._dedupe_candidates( - [ - MemoryCandidate(candidate_id="cand-template-a", user_id="u", memory_type=MemoryType.TEMPLATE, key="template.generic.x", content="a"), - MemoryCandidate(candidate_id="cand-template-b", user_id="u", memory_type=MemoryType.TEMPLATE, key="template.generic.x", content="b"), - ] - ) - assert len(deduped) == 1 - - -def test_knowledge_templates_are_isolated_per_user_and_have_stable_keys() -> None: - events = [ - _event( - "alice-1", - user_id="alice", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="exporter", - content="批量导出订单到 Excel", - input_payload={"operation": "batch export", "source": "orders_001.csv"}, - output_payload={"file": "orders_001.xlsx"}, - ), - _event( - "alice-2", - user_id="alice", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="exporter", - content="批量导出订单到 Excel", - input_payload={"operation": "batch export", "source": "orders_002.csv"}, - output_payload={"file": "orders_002.xlsx"}, - ), - _event( - "bob-1", - user_id="bob", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="exporter", - content="批量导出订单到 Excel", - input_payload={"operation": "batch export", "source": "orders_003.csv"}, - output_payload={"file": "orders_003.xlsx"}, - ), - _event( - "bob-2", - user_id="bob", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="exporter", - content="批量导出订单到 Excel", - input_payload={"operation": "batch export", "source": "orders_004.csv"}, - output_payload={"file": "orders_004.xlsx"}, - ), - ] - - candidates = KnowledgeExtractor.extract_templates(list(reversed(events))) - assert {candidate.user_id for candidate in candidates} == {"alice", "bob"} - assert all(all(event_id.startswith(candidate.user_id) for event_id in candidate.source_events) for candidate in candidates) - - repeated = KnowledgeExtractor.extract_templates(events) - assert {(candidate.user_id, candidate.key, candidate.candidate_id) for candidate in candidates} == { - (candidate.user_id, candidate.key, candidate.candidate_id) for candidate in repeated - } + assert llm_client.calls[0]["request"]["mode"] == "knowledge_from_tool_result" + assert candidates[0].memory_type is MemoryType.KNOWLEDGE + assert candidates[0].key == "knowledge.tool_case.batch_export" -def test_preference_explicit_allows_modifiers_between_cue_and_format() -> None: - event = _event( - "evt-pref-flex-1", - content="以后导出都用 PDF 格式;下次生成报告默认保存成 Markdown。", - ) +def test_knowledge_template_extraction_uses_llm(llm_client: RoutingFakeLLMClient) -> None: + events = [_event("evt-template-1", content="合并 a.csv 和 b.csv"), _event("evt-template-2", content="再次合并文件")] - candidates = PreferenceExtractor.extract_from_conversation(event) - keys = {candidate.key for candidate in candidates} + candidates = KnowledgeExtractor.extract_templates(events) - assert "preference.output_format.pdf" in keys - assert "preference.output_format.markdown" in keys + assert llm_client.calls[0]["request"]["mode"] == "template_from_session" + assert candidates[0].memory_type is MemoryType.TEMPLATE + assert candidates[0].key == "template.data_processing.merge_files" -def test_preference_ignores_empty_and_truncates_extreme_noise() -> None: - assert PreferenceExtractor.extract_explicit_preference("") == [] - noisy = "x" * 15000 + " 以后都用 PDF 输出" - assert PreferenceExtractor.extract_explicit_preference(noisy) == [] +def test_knowledge_deduplication_is_done_after_llm_output() -> None: + duplicate = _payload("knowledge", "tool_case", "batch_export", "工具结果可复用为批量导出知识") + duplicate["candidates"].append(dict(duplicate["candidates"][0])) + class DuplicateClient: + def complete_json(self, prompt: str, schema: dict[str, Any]) -> dict[str, Any]: + return duplicate -def test_knowledge_tool_result_redacts_sensitive_values() -> None: - event = _event( - "evt-secret-1", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="deploy", - output_payload={ - "status": "success", - "api_key": "sk-live-1234567890abcdef", - "message": "done with password=super-secret token=abcdef1234567890abcdef123456", - }, - metadata={"authorization": "Bearer abcdef1234567890abcdef123456"}, - ) + set_default_llm_client(DuplicateClient()) + try: + candidates = KnowledgeExtractor.extract_from_conversation(_event("evt-dup", content="导出成功")) + finally: + clear_default_llm_client() - candidates = KnowledgeExtractor.extract_from_tool_result(event) - rendered = str([candidate.to_dict() for candidate in candidates]) - - assert "sk-live" not in rendered - assert "super-secret" not in rendered - assert "abcdef1234567890abcdef123456" not in rendered - assert "" in rendered + assert len(candidates) == 1 diff --git a/tests/test_ingestion_to_extractors.py b/tests/test_ingestion_to_extractors.py index 3a538ce..b943585 100644 --- a/tests/test_ingestion_to_extractors.py +++ b/tests/test_ingestion_to_extractors.py @@ -1,105 +1,116 @@ from __future__ import annotations -from ingestion.adapter import raw_event_to_memory_event -from ingestion.collector import create_raw_event +import json +from datetime import datetime, timezone +from typing import Any + +import pytest + +from core.constants import EventType, MemoryType, Scene +from core.models import MemoryEvent from extractors.environment_extractor import EnvironmentExtractor from extractors.knowledge_extractor import KnowledgeExtractor +from extractors.llm_memory_extractor import clear_default_llm_client, set_default_llm_client from extractors.preference_extractor import PreferenceExtractor from extractors.tool_extractor import ToolExtractor -def _memory_event(payload: dict): - return raw_event_to_memory_event(create_raw_event(payload)) +class IntegrationFakeLLMClient: + def __init__(self) -> None: + self.calls: list[dict[str, Any]] = [] + + def complete_json(self, prompt: str, schema: dict[str, Any]) -> dict[str, Any]: + request = json.loads(prompt) + self.calls.append({"request": request, "schema": schema}) + mode = request["mode"] + if mode.startswith("preference"): + return _payload("preference", "output_format", "markdown", "用户偏好 Markdown") + if mode.startswith("knowledge"): + return _payload("knowledge", "tool_case", "export_pdf", "PDF 导出工具结果可复用") + if mode.startswith("tool"): + return _payload("tool", "experience", "export_pdf", "export_pdf 工具适合文档导出") + if mode.startswith("environment"): + return _payload("environment", "path", "documents", "常用文档目录是 /root/docs") + return {"candidates": []} + +def _payload(memory_type: str, category: str, value: str, content: str) -> dict[str, Any]: + return { + "candidates": [ + { + "is_memory_worthy": True, + "is_long_term": True, + "memory_type": memory_type, + "category": category, + "value": value, + "scope": "integration", + "content": content, + "confidence": 0.84, + "evidence": "fake integration evidence", + "reason": "fake integration reason", + "sensitivity": "none", + } + ] + } -def test_ingestion_conversation_event_feeds_preference_extractor() -> None: - event = _memory_event( - { - "event_id": "raw-pref-1", - "user_id": "user-a", - "session_id": "session-a", - "task_id": "task-a", - "event_type": "conversation", - "scenario": "office", - "timestamp": "2026-06-27T10:00:00+08:00", - "content": "以后都用 Markdown 输出,回答尽量简洁。", - } + +@pytest.fixture() +def integration_client() -> IntegrationFakeLLMClient: + client = IntegrationFakeLLMClient() + set_default_llm_client(client) + try: + yield client + finally: + clear_default_llm_client() + + +def _event(event_id: str, event_type: EventType, **kwargs: Any) -> MemoryEvent: + return MemoryEvent( + event_id=event_id, + raw_event_id=f"raw-{event_id}", + user_id="u-ingest", + session_id="s-ingest", + task_id="t-ingest", + event_type=event_type, + scenario=Scene.OFFICE, + source=event_type.value, + actor="user", + timestamp=datetime(2026, 7, 5, 10, 0, tzinfo=timezone.utc), + **kwargs, ) + +def test_ingestion_output_can_drive_preference_llm_extractor(integration_client: IntegrationFakeLLMClient) -> None: + event = _event(EventType.CONVERSATION.value, EventType.CONVERSATION, content="以后文档用 Markdown 输出") + candidates = PreferenceExtractor.extract_from_conversation(event) - keys = {candidate.key for candidate in candidates} - - assert "preference.output_format.markdown" in keys - assert all(candidate.user_id == "user-a" for candidate in candidates) - assert all("raw-pref-1" in candidate.source_events for candidate in candidates) - - -def test_ingestion_tool_events_feed_knowledge_and_tool_extractors() -> None: - call_event = _memory_event( - { - "event_id": "raw-tool-call-1", - "user_id": "user-a", - "session_id": "session-a", - "task_id": "task-a", - "event_type": "tool_call", - "scenario": "office", - "timestamp": "2026-06-27T10:01:00+08:00", - "tool_name": "wps_export", - "input": {"file": "report.docx", "format": "pdf"}, - "call_id": "call-1", - } - ) - result_event = _memory_event( - { - "event_id": "raw-tool-result-1", - "user_id": "user-a", - "session_id": "session-a", - "task_id": "task-a", - "event_type": "tool_result", - "scenario": "office", - "timestamp": "2026-06-27T10:01:03+08:00", - "tool_name": "wps_export", - "output": {"file": "report.pdf", "status": "success"}, - "success": True, - "duration_ms": 3000, - "call_id": "call-1", - } - ) - knowledge = KnowledgeExtractor.extract_from_tool_result(result_event) - tool_patterns = ToolExtractor.extract_tool_pattern([call_event, result_event]) - - assert knowledge - assert any(candidate.key == "knowledge.tool_case.batch_export.wps_export" for candidate in knowledge) - assert len(tool_patterns) == 1 - assert tool_patterns[0].metadata["success_rate"] == 1.0 - assert tool_patterns[0].metadata["success_count"] == 1 - - -def test_ingestion_system_context_event_feeds_environment_extractor() -> None: - event = _memory_event( - { - "event_id": "raw-env-1", - "user_id": "user-a", - "session_id": "session-a", - "task_id": "task-a", - "event_type": "system_context", - "scenario": "system", - "timestamp": "2026-06-27T10:02:00+08:00", - "downloads": "C:\\Users\\alice\\Downloads", - "documents": "C:\\Users\\alice\\Documents", - "locale": "zh_CN", - "installed_software": ["WPS Office", "Python"], - "os_version": "Kylin Desktop V11", - } + assert integration_client.calls[0]["request"]["events"][0]["event_id"] == EventType.CONVERSATION.value + assert candidates[0].memory_type is MemoryType.PREFERENCE + + +def test_tool_result_event_can_drive_knowledge_and_tool_llm_extractors(integration_client: IntegrationFakeLLMClient) -> None: + event = _event( + "tool-result", + EventType.TOOL_RESULT, + tool_name="export_pdf", + output={"status": "success", "file": "report.pdf"}, + success=True, ) - candidates = EnvironmentExtractor.extract_from_tool_output(event.metadata) - keys = {candidate.key for candidate in candidates} + knowledge = KnowledgeExtractor.extract_from_tool_result(event) + tool_memory = ToolExtractor.extract_tool_pattern([event]) + + assert knowledge[0].memory_type is MemoryType.KNOWLEDGE + assert tool_memory[0].memory_type is MemoryType.TOOL + assert [call["request"]["mode"] for call in integration_client.calls] == [ + "knowledge_from_tool_result", + "tool_pattern", + ] + + +def test_environment_metadata_can_drive_environment_llm_extractor(integration_client: IntegrationFakeLLMClient) -> None: + candidates = EnvironmentExtractor.extract_from_tool_output({"user_id": "u-ingest", "documents": "/root/docs"}) - assert "environment.path.downloads" in keys - assert "environment.path.documents" in keys - assert "environment.locale.language" in keys - assert "environment.locale.region" in keys - assert any(candidate.key.startswith("environment.software.") for candidate in candidates) - assert any(candidate.metadata.get("path") == "~\\Downloads" for candidate in candidates) + assert candidates[0].memory_type is MemoryType.ENVIRONMENT + assert integration_client.calls[0]["request"]["events"][0]["event_type"] == "system_context" diff --git a/tests/test_llm_memory_extractor.py b/tests/test_llm_memory_extractor.py new file mode 100644 index 0000000..24b44d0 --- /dev/null +++ b/tests/test_llm_memory_extractor.py @@ -0,0 +1,320 @@ +from __future__ import annotations + +import json +from datetime import datetime, timezone +from typing import Any + +from core.constants import EventType, MemoryType, Scene +from core.models import MemoryCandidate, MemoryEvent +from extractors.llm_memory_extractor import ( + CandidateMerger, + CandidateValidator, + HybridMemoryExtractor, + LLMMemoryExtractor, + LLMWorkflowBoundaryExtractor, + clear_default_llm_client, + should_call_llm, +) + + +class FakeLLMClient: + def __init__(self, payload: dict[str, Any] | list[dict[str, Any]]) -> None: + self.payload = payload + self.calls: list[dict[str, Any]] = [] + + def complete_json(self, prompt: str, schema: dict[str, Any]) -> Any: + self.calls.append({"prompt": prompt, "schema": schema, "request": json.loads(prompt)}) + return self.payload + + +def _event( + event_id: str, + *, + event_type: EventType = EventType.CONVERSATION, + source: str = "conversation", + content: str | None = None, + tool_name: str | None = None, + input_payload: dict[str, Any] | None = None, + output_payload: dict[str, Any] | None = None, + metadata: dict[str, Any] | None = None, + success: bool | None = True, +) -> MemoryEvent: + return MemoryEvent( + event_id=event_id, + raw_event_id=f"raw-{event_id}", + user_id="user-llm", + session_id="session-llm", + task_id="task-llm", + event_type=event_type, + scenario=Scene.OFFICE, + source=source, + actor="user" if event_type is EventType.CONVERSATION else "agent", + content=content, + tool_name=tool_name, + input=input_payload or {}, + output=output_payload or {}, + metadata=metadata or {}, + success=success, + timestamp=datetime(2026, 7, 5, 10, 0, tzinfo=timezone.utc), + ) + + +def _payload( + *, + memory_type: str = "preference", + category: str = "response_order", + value: str = "conclusion_first", + content: str = "用户偏好报告类内容先给结论再展开", + confidence: float = 0.84, + is_memory_worthy: bool = True, + is_long_term: bool = True, +) -> dict[str, Any]: + return { + "candidates": [ + { + "is_memory_worthy": is_memory_worthy, + "is_long_term": is_long_term, + "memory_type": memory_type, + "category": category, + "value": value, + "scope": "report_writing", + "content": content, + "confidence": confidence, + "evidence": "这种报告以后先给结论再展开。", + "reason": "LLM 判断这是可复用的长期写作偏好。", + "sensitivity": "none", + } + ] + } + + +def test_llm_prompt_contains_sanitized_input_dataset_before_model_output() -> None: + event = _event( + "evt-before-after-1", + content="以后报告先给结论。api_key=sk-secret-1234567890", + metadata={"email": "alice@example.com", "api_key": "sk-secret-1234567890"}, + ) + client = FakeLLMClient(_payload()) + + candidates = LLMMemoryExtractor.extract_event(event, client, mode="conversation") + + assert candidates[0].key == "preference.response_order.conclusion_first" + prompt = client.calls[0]["prompt"] + request = client.calls[0]["request"] + assert request["mode"] == "conversation" + assert request["events"][0]["event_id"] == "evt-before-after-1" + assert "sk-secret-1234567890" not in prompt + assert "alice@example.com" not in prompt + assert "[REDACTED_SECRET]" in prompt + assert "[REDACTED_EMAIL]" in prompt + + +def test_hybrid_facade_always_uses_llm_when_client_is_supplied() -> None: + event = _event("evt-llm-2", content="以后都用 Markdown 输出。") + client = FakeLLMClient(_payload(category="output_format", value="markdown", content="用户偏好 Markdown 输出")) + + candidates = HybridMemoryExtractor.extract_from_conversation(event, client) + + assert client.calls + assert [candidate.key for candidate in candidates] == ["preference.output_format.markdown"] + assert candidates[0].source == "llm_extracted" + + +def test_should_call_llm_has_no_rule_or_length_gate() -> None: + assert should_call_llm(_event("evt-empty", content="")) is False + assert should_call_llm(_event("evt-short", content="好")) is True + assert should_call_llm(_event("evt-tool", event_type=EventType.TOOL_RESULT, tool_name="export", output_payload={"ok": True})) is True + + +def test_validator_rejects_temporary_and_credential_like_candidates() -> None: + event = _event("evt-llm-4", content="这次临时用 PDF,api_key=sk-1234567890abcdef。") + client = FakeLLMClient( + { + "candidates": [ + { + "is_memory_worthy": True, + "is_long_term": False, + "memory_type": "preference", + "category": "output_format", + "value": "pdf", + "scope": "current_task", + "content": "用户这次临时使用 PDF", + "confidence": 0.9, + "evidence": "这次临时用 PDF", + "reason": "当前任务约束,不是长期偏好。", + "sensitivity": "none", + }, + { + "is_memory_worthy": True, + "is_long_term": True, + "memory_type": "knowledge", + "category": "credential", + "value": "api_key", + "scope": "tool_runtime", + "content": "用户 API_KEY 是 sk-1234567890abcdef", + "confidence": 0.88, + "evidence": "api_key=sk-1234567890abcdef", + "reason": "包含凭据,不应进入长期记忆。", + "sensitivity": "api_key", + }, + ] + } + ) + + candidates = LLMMemoryExtractor.extract_event(event, client) + + assert candidates == [] + assert "sk-1234567890abcdef" not in client.calls[0]["prompt"] + + +def test_validator_redacts_contact_information_without_changing_extraction_logic() -> None: + candidate = MemoryCandidate( + candidate_id="contact-1", + user_id="user-llm", + memory_type=MemoryType.SAFETY, + key="safety.contact.redaction", + content="不要长期保存 13812345678 和 alice@example.com", + scenario=Scene.OFFICE, + confidence=0.91, + source="llm_extracted", + metadata={"evidence": "用户要求不要保存联系方式"}, + ) + + checked = CandidateValidator().validate(candidate) + + assert checked is not None + assert "[REDACTED_PHONE]" in checked.content + assert "[REDACTED_EMAIL]" in checked.content + assert checked.confidence == 0.91 + assert checked.metadata["sensitive_redacted"] is True + + +def test_merger_deduplicates_and_marks_conflicts_without_rule_boosting() -> None: + left = MemoryCandidate( + candidate_id="pref-1", + user_id="user-llm", + memory_type=MemoryType.PREFERENCE, + key="preference.output_format.markdown", + content="用户偏好 Markdown", + scenario=Scene.OFFICE, + confidence=0.8, + source="llm_extracted", + metadata={"evidence": "以后用 Markdown"}, + ) + better_duplicate = MemoryCandidate( + candidate_id="pref-1b", + user_id="user-llm", + memory_type=MemoryType.PREFERENCE, + key="preference.output_format.markdown", + content="用户强烈偏好 Markdown", + scenario=Scene.OFFICE, + confidence=0.9, + source="llm_extracted", + metadata={"evidence": "默认用 Markdown"}, + ) + conflict = MemoryCandidate( + candidate_id="pref-2", + user_id="user-llm", + memory_type=MemoryType.PREFERENCE, + key="preference.output_format.pdf", + content="用户偏好 PDF", + scenario=Scene.OFFICE, + confidence=0.83, + source="llm_extracted", + metadata={"evidence": "以后用 PDF"}, + ) + + merged = CandidateMerger.merge([left], [better_duplicate, conflict]) + + assert len(merged) == 2 + markdown = next(candidate for candidate in merged if candidate.key.endswith("markdown")) + assert markdown.content == "用户强烈偏好 Markdown" + for candidate in merged: + assert candidate.metadata["possible_conflict_keys"] + + +def test_workflow_boundaries_are_model_output_not_local_markers() -> None: + events = [ + _event("evt-0", content="开始处理月报"), + _event("evt-1", event_type=EventType.TOOL_RESULT, tool_name="read_sheet", output_payload={"rows": 20}), + _event("evt-2", event_type=EventType.TOOL_RESULT, tool_name="export_pdf", output_payload={"file": "report.pdf"}), + ] + client = FakeLLMClient({"boundaries": [{"start": 0, "end": 2, "reason": "同一月报流程"}]}) + + assert LLMWorkflowBoundaryExtractor.detect_boundaries(events, client) == [(0, 2)] + + +def test_llm_parser_handles_empty_events_inactive_events_invalid_types_and_list_shape() -> None: + client = FakeLLMClient([]) + assert LLMMemoryExtractor.extract_events([], client) == [] + assert LLMMemoryExtractor.extract_event(_event("evt-inactive"), client) == [] + + event = _event("evt-list", content="以后报告按模板整理") + client = FakeLLMClient( + [ + { + "is_memory_worthy": True, + "is_long_term": True, + "memory_type": "not_a_type", + "category": "x", + "value": "x", + "content": "invalid", + "confidence": 1, + "evidence": "invalid", + "reason": "invalid", + }, + { + "is_memory_worthy": True, + "is_long_term": True, + "memory_type": "preference", + "category": "report_style", + "value": "template", + "content": "", + "confidence": 1, + "evidence": "empty", + "reason": "empty", + }, + { + "is_memory_worthy": True, + "is_long_term": True, + "memory_type": "preference", + "category": "report_style", + "value": "template", + "content": "用户偏好按模板整理报告", + "confidence": "bad-score", + "evidence": "以后报告按模板整理", + "reason": "LLM 判断为长期偏好", + }, + ] + ) + + candidates = LLMMemoryExtractor.extract_event(event, client) + + assert len(candidates) == 1 + assert candidates[0].key == "preference.report_style.template" + assert candidates[0].confidence == 0.0 + + +def test_hybrid_facade_without_any_client_returns_empty_lists() -> None: + clear_default_llm_client() + event = _event("evt-no-client", content="以后先给结论") + + assert HybridMemoryExtractor.extract_from_conversation(event) == [] + assert HybridMemoryExtractor.extract_from_tool_result(event) == [] + assert HybridMemoryExtractor.extract_from_session([event]) == [] + + +def test_boundary_parser_ignores_invalid_model_items() -> None: + events = [_event("evt-boundary-0", content="开始"), _event("evt-boundary-1", content="结束")] + client = FakeLLMClient( + { + "boundaries": [ + {"start": 0, "end": 1, "reason": "valid"}, + {"start": 1, "end": 5, "reason": "out of range"}, + {"start": "bad", "end": 1, "reason": "bad"}, + "bad", + ] + } + ) + + assert LLMWorkflowBoundaryExtractor.detect_boundaries(events, client) == [(0, 1)] diff --git a/tests/test_memory_admission.py b/tests/test_memory_admission.py new file mode 100644 index 0000000..867f1a3 --- /dev/null +++ b/tests/test_memory_admission.py @@ -0,0 +1,368 @@ +"""End-to-end tests for MemoryCandidate -> MemoryRecord admission.""" + +from __future__ import annotations + +import threading +from concurrent.futures import ThreadPoolExecutor +from datetime import datetime, timezone + +import pytest + +from core.constants import EventType, MemoryStatus, MemoryType, Scene +from core.models import MemoryCandidate, MemoryEvent +from extractors.llm_memory_extractor import LLMMemoryExtractor +from memory.admission import MemoryAdmissionService, SQLiteAdmissionRepository +from memory.conflict_resolver import ConflictAction, LLMConflictResolver +from memory.lifecycle_state import InvalidMemoryTransition + + +def candidate(**overrides) -> MemoryCandidate: + values = { + "candidate_id": "cand_pdf", + "user_id": "user-1", + "memory_type": MemoryType.PREFERENCE, + "key": "preference.export.format", + "content": "用户偏好使用 PDF 导出", + "scenario": Scene.OFFICE, + "confidence": 0.91, + "source": "llm_extracted", + "source_events": ["evt-1"], + "tags": ["preference", "export"], + "metadata": {"evidence": "以后都导出 PDF"}, + } + values.update(overrides) + return MemoryCandidate(**values) + + +def response(action: str, **overrides) -> dict: + value = { + "action": action, + "reason": f"model selected {action}", + "target_memory_ids": [], + "final_content": "用户偏好使用 PDF 导出", + "requires_human_review": False, + "decision_confidence": 0.96, + } + value.update(overrides) + return value + + +class QueueLLM: + def __init__(self, *responses): + self.responses = list(responses) + self.prompts: list[str] = [] + self._lock = threading.Lock() + + def complete_json(self, prompt, schema): + with self._lock: + self.prompts.append(prompt) + item = self.responses.pop(0) + if isinstance(item, Exception): + raise item + return item + + +class ConcurrentCreateLLM: + def __init__(self): + self.barrier = threading.Barrier(2) + self.calls = 0 + self._lock = threading.Lock() + + def complete_json(self, prompt, schema): + with self._lock: + self.calls += 1 + self.barrier.wait(timeout=5) + return response("create") + + +@pytest.fixture +def repository(tmp_path): + return SQLiteAdmissionRepository(str(tmp_path / "memory.db")) + + +def service(repository, llm) -> MemoryAdmissionService: + return MemoryAdmissionService(repository, LLMConflictResolver(llm)) + + +def test_create_promotes_candidate_to_active_record(repository): + llm = QueueLLM(response("create")) + result = service(repository, llm).admit(candidate()) + + assert result.action is ConflictAction.CREATE + assert result.status is MemoryStatus.ACTIVE + assert result.record is not None + assert result.record.version == 1 + assert result.record.content == "用户偏好使用 PDF 导出" + stored = repository.list_all() + assert len(stored) == 1 + assert stored[0].memory_id == result.record.memory_id + assert stored[0].status is MemoryStatus.ACTIVE + + +def test_same_candidate_retry_is_idempotent_without_second_llm_call(repository): + llm = QueueLLM(response("create")) + admission = service(repository, llm) + + first = admission.admit(candidate()) + second = admission.admit(candidate()) + + assert first.record.memory_id == second.record.memory_id + assert second.action is ConflictAction.DUPLICATE + assert second.reason == "idempotent_candidate_retry" + assert len(llm.prompts) == 1 + assert len(repository.list_all()) == 1 + + +def test_semantic_duplicate_reuses_existing_record(repository): + first_llm = QueueLLM(response("create")) + first = service(repository, first_llm).admit(candidate()) + duplicate_candidate = candidate( + candidate_id="cand_pdf_rephrased", + content="导出文件时优先采用 PDF 格式", + ) + duplicate_llm = QueueLLM( + response( + "duplicate", + target_memory_ids=[first.record.memory_id], + final_content="", + ) + ) + + result = service(repository, duplicate_llm).admit(duplicate_candidate) + + assert result.action is ConflictAction.DUPLICATE + assert result.record.memory_id == first.record.memory_id + assert len(repository.list_all()) == 1 + + +def test_replace_supersedes_old_record_and_increments_version(repository): + first = service(repository, QueueLLM(response("create"))).admit(candidate()) + newer = candidate( + candidate_id="cand_word", + content="用户现在偏好使用 Word 导出", + ) + replace_llm = QueueLLM( + response( + "replace", + target_memory_ids=[first.record.memory_id], + final_content=newer.content, + ) + ) + + result = service(repository, replace_llm).admit(newer) + records = repository.list_all() + + assert result.action is ConflictAction.REPLACE + assert result.record.version == 2 + assert result.record.supersedes == first.record.memory_id + assert [record.status for record in records] == [ + MemoryStatus.SUPERSEDED, + MemoryStatus.ACTIVE, + ] + + +def test_merge_creates_canonical_new_version(repository): + first = service(repository, QueueLLM(response("create"))).admit(candidate()) + supplement = candidate( + candidate_id="cand_pdf_quality", + content="PDF 导出时使用高质量模式", + ) + merged_content = "用户偏好使用 PDF 导出,并启用高质量模式" + merge_llm = QueueLLM( + response( + "merge", + target_memory_ids=[first.record.memory_id], + final_content=merged_content, + ) + ) + + result = service(repository, merge_llm).admit(supplement) + + assert result.action is ConflictAction.MERGE + assert result.record.content == merged_content + assert result.record.version == 2 + assert repository.get(first.record.memory_id).status is MemoryStatus.SUPERSEDED + + +def test_different_scenario_memories_can_coexist(repository): + office = service(repository, QueueLLM(response("create"))).admit(candidate()) + coding_candidate = candidate( + candidate_id="cand_markdown", + scenario=Scene.CODING, + content="编程文档偏好 Markdown", + ) + coexist_llm = QueueLLM( + response("coexist", final_content=coding_candidate.content) + ) + + coding = service(repository, coexist_llm).admit(coding_candidate) + + assert office.record.status is MemoryStatus.ACTIVE + assert coding.action is ConflictAction.COEXIST + assert coding.record.version == 1 + assert len( + repository.list_related( + coding_candidate, + statuses=(MemoryStatus.ACTIVE,), + ) + ) == 2 + + +def test_invalid_model_response_is_persisted_as_pending(repository): + llm = QueueLLM({"action": "unsafe_unknown_action"}) + result = service(repository, llm).admit(candidate()) + + assert result.action is ConflictAction.PENDING + assert result.status is MemoryStatus.PENDING + assert result.record.status is MemoryStatus.PENDING + assert repository.list_all()[0].status is MemoryStatus.PENDING + + +def test_pending_candidate_can_be_promoted_after_model_recovers(repository): + admission = service( + repository, + QueueLLM({"action": "invalid"}, response("create")), + ) + + pending = admission.admit(candidate()) + promoted = admission.admit(candidate()) + + assert pending.status is MemoryStatus.PENDING + assert promoted.status is MemoryStatus.ACTIVE + assert promoted.record.memory_id == pending.record.memory_id + assert len(repository.list_all()) == 1 + + +def test_model_reject_does_not_create_formal_memory(repository): + reject_llm = QueueLLM( + response("reject", final_content="", reason="not durable memory") + ) + result = service(repository, reject_llm).admit(candidate()) + + assert result.action is ConflictAction.REJECT + assert result.status is MemoryStatus.REJECTED + assert result.record is None + assert repository.list_all() == [] + + +def test_sensitive_candidate_is_rejected_before_llm(repository): + llm = QueueLLM(response("create")) + result = service(repository, llm).admit( + candidate(content="password=super-secret-value") + ) + + assert result.action is ConflictAction.REJECT + assert result.reason == "candidate_failed_security_or_completeness_validation" + assert llm.prompts == [] + assert repository.list_all() == [] + + +def test_llm_outage_fails_closed_without_database_mutation(repository): + result = service(repository, QueueLLM(RuntimeError("offline"))).admit(candidate()) + + assert result.action is ConflictAction.PENDING + assert result.status is MemoryStatus.PENDING + assert result.record is None + assert result.reason == "llm_unavailable:RuntimeError" + assert repository.list_all() == [] + + +def test_lifecycle_transitions_are_validated(repository): + admission = service(repository, QueueLLM(response("create"))) + created = admission.admit(candidate()).record + + archived = admission.transition(created.memory_id, MemoryStatus.ARCHIVED) + restored = admission.transition(created.memory_id, MemoryStatus.ACTIVE) + + assert archived.status is MemoryStatus.ARCHIVED + assert restored.status is MemoryStatus.ACTIVE + with pytest.raises(InvalidMemoryTransition): + admission.transition(created.memory_id, MemoryStatus.REJECTED) + + +def test_concurrent_same_candidate_creates_only_one_record(repository): + llm = ConcurrentCreateLLM() + admission = service(repository, llm) + + with ThreadPoolExecutor(max_workers=2) as executor: + results = list(executor.map(admission.admit, [candidate(), candidate()])) + + assert {result.action for result in results} == { + ConflictAction.CREATE, + ConflictAction.DUPLICATE, + } + assert len(repository.list_all()) == 1 + assert llm.calls == 2 + + +def test_replace_transaction_rolls_back_when_insert_fails(repository, monkeypatch): + first = service(repository, QueueLLM(response("create"))).admit(candidate()) + newer = candidate( + candidate_id="cand-word-rollback", + content="用户偏好使用 Word 导出", + ) + replace_llm = QueueLLM( + response( + "replace", + target_memory_ids=[first.record.memory_id], + final_content=newer.content, + ) + ) + + def fail_insert(*args, **kwargs): + raise RuntimeError("simulated storage failure") + + monkeypatch.setattr(repository, "_insert_record", fail_insert) + with pytest.raises(RuntimeError, match="simulated storage failure"): + service(repository, replace_llm).admit(newer) + + records = repository.list_all() + assert len(records) == 1 + assert records[0].memory_id == first.record.memory_id + assert records[0].status is MemoryStatus.ACTIVE + + +def test_memory_event_can_flow_through_llm_extraction_into_formal_storage(repository): + event = MemoryEvent( + event_id="event-e2e", + raw_event_id="raw-e2e", + user_id="user-1", + session_id="session-e2e", + task_id="task-e2e", + event_type=EventType.CONVERSATION, + scenario=Scene.OFFICE, + source="conversation", + actor="user", + content="以后办公文件统一使用 PDF 导出", + timestamp=datetime(2026, 7, 31, tzinfo=timezone.utc), + ) + extraction_llm = QueueLLM( + { + "candidates": [ + { + "is_memory_worthy": True, + "is_long_term": True, + "memory_type": "preference", + "category": "export", + "value": "pdf", + "scope": "office", + "content": "用户偏好办公文件使用 PDF 导出", + "confidence": 0.91, + "evidence": event.content, + "reason": "用户明确表达了长期偏好", + "sensitivity": "none", + } + ] + } + ) + extracted = LLMMemoryExtractor.extract_event(event, extraction_llm) + admission_llm = QueueLLM( + response("create", final_content=extracted[0].content) + ) + + result = service(repository, admission_llm).admit(extracted[0]) + + assert len(extracted) == 1 + assert result.status is MemoryStatus.ACTIVE + assert result.record.memory_type is MemoryType.PREFERENCE + assert result.record.source_events == ["event-e2e"] diff --git a/tests/test_tool_extractor.py b/tests/test_tool_extractor.py index 4300a33..a1ffb5c 100644 --- a/tests/test_tool_extractor.py +++ b/tests/test_tool_extractor.py @@ -1,201 +1,116 @@ from __future__ import annotations -from datetime import datetime, timedelta, timezone +import json +from datetime import datetime, timezone +from typing import Any + +import pytest from core.constants import EventType, MemoryType, Scene from core.models import MemoryEvent +from extractors.llm_memory_extractor import clear_default_llm_client, set_default_llm_client from extractors.tool_extractor import ToolExtractor -def _event( - event_id: str, - *, - user_id: str = "alice", - session_id: str = "session-1", - task_id: str = "task-1", - event_type: EventType, - tool_name: str = "backup", - input_payload: dict | None = None, - output_payload: dict | None = None, - metadata: dict | None = None, - success: bool | None = None, - offset_seconds: int = 0, -) -> MemoryEvent: +class ToolFakeLLMClient: + def __init__(self) -> None: + self.calls: list[dict[str, Any]] = [] + + def complete_json(self, prompt: str, schema: dict[str, Any]) -> dict[str, Any]: + request = json.loads(prompt) + self.calls.append({"request": request, "schema": schema}) + return { + "candidates": [ + { + "is_memory_worthy": True, + "is_long_term": True, + "memory_type": "tool", + "category": "success_rate", + "value": "backup", + "scope": "tool_usage", + "content": "backup 工具成功率约 0.67,适合常规备份任务。", + "confidence": 0.82, + "evidence": "LLM reviewed tool calls and results.", + "reason": "LLM 从工具调用轨迹中总结工具经验。", + "sensitivity": "none", + } + ] + } + + +@pytest.fixture() +def tool_client() -> ToolFakeLLMClient: + client = ToolFakeLLMClient() + set_default_llm_client(client) + try: + yield client + finally: + clear_default_llm_client() + + +def _event(event_id: str, event_type: EventType, tool_name: str, success: bool | None = None) -> MemoryEvent: return MemoryEvent( event_id=event_id, raw_event_id=f"raw-{event_id}", - user_id=user_id, - session_id=session_id, - task_id=task_id, + user_id="u-tool", + session_id="s-tool", + task_id="t-tool", event_type=event_type, scenario=Scene.SYSTEM, - source="tool", + source=event_type.value, + actor="agent", tool_name=tool_name, - input=input_payload or {}, - output=output_payload or {}, - metadata=metadata or {}, success=success, - timestamp=datetime(2026, 6, 25, tzinfo=timezone.utc) + timedelta(seconds=offset_seconds), + timestamp=datetime(2026, 7, 5, 10, 0, tzinfo=timezone.utc), ) -def test_tool_pattern_statistics_are_exact_and_user_isolated() -> None: - events = [ - _event( - "alice-call-1", - event_type=EventType.TOOL_CALL, - input_payload={"mode": "full", "target": "docs", "api_key": "must-not-leak"}, - metadata={"call_id": "a-1"}, - offset_seconds=1, - ), - _event( - "alice-result-1", - event_type=EventType.TOOL_RESULT, - output_payload={"status": "ok"}, - metadata={"call_id": "a-1", "duration_ms": 120}, - success=True, - offset_seconds=2, - ), - _event( - "alice-call-2", - event_type=EventType.TOOL_CALL, - input_payload={"mode": "incremental", "target": "docs"}, - metadata={"call_id": "a-2"}, - offset_seconds=3, - ), - _event( - "alice-result-2", - event_type=EventType.TOOL_RESULT, - output_payload={"error_type": "timeout", "message": "remote timeout after 30 seconds"}, - metadata={"call_id": "a-2", "duration_ms": 300}, - success=False, - offset_seconds=4, - ), - _event( - "alice-call-3", - event_type=EventType.TOOL_CALL, - input_payload={"mode": "full", "target": "documents"}, - metadata={"call_id": "a-3"}, - offset_seconds=5, - ), - _event( - "bob-call-1", - user_id="bob", - event_type=EventType.TOOL_CALL, - input_payload={"mode": "full"}, - metadata={"call_id": "b-1"}, - offset_seconds=1, - ), - _event( - "bob-result-1", - user_id="bob", - event_type=EventType.TOOL_RESULT, - output_payload={"status": "ok"}, - metadata={"call_id": "b-1", "duration_ms": 80}, - success=True, - offset_seconds=2, - ), - ] - - candidates = ToolExtractor.extract_tool_pattern(list(reversed(events))) - assert {candidate.user_id for candidate in candidates} == {"alice", "bob"} - - alice = next(candidate for candidate in candidates if candidate.user_id == "alice") - assert alice.memory_type is MemoryType.TOOL - assert alice.metadata["success_count"] == 1 - assert alice.metadata["failure_count"] == 1 - assert alice.metadata["total_count"] == 2 - assert alice.metadata["observed_count"] == 3 - assert alice.metadata["unknown_count"] == 1 - assert alice.metadata["success_rate"] == 0.5 - assert alice.metadata["average_response_ms"] == 210.0 - assert alice.metadata["common_failure_reasons"][0]["reason"].startswith("error_type: timeout") - assert "api_key" not in str(alice.metadata["common_parameter_combinations"]) - assert all(event_id.startswith("alice") for event_id in alice.source_events) - - assert ToolExtractor.calculate_tool_success_rate("backup", events) == 2 / 3 - - -def test_tool_success_rate_pairs_calls_and_results_by_time_and_correlation() -> None: - events = [ - _event( - "result-2", - event_type=EventType.TOOL_RESULT, - metadata={"call_id": "two", "duration_ms": 20}, - output_payload={"status": "failed", "error": "network error"}, - offset_seconds=4, - ), - _event( - "call-1", - event_type=EventType.TOOL_CALL, - metadata={"call_id": "one"}, - input_payload={"format": "json"}, - offset_seconds=1, - ), - _event( - "result-1", - event_type=EventType.TOOL_RESULT, - metadata={"call_id": "one", "duration_ms": 10}, - output_payload={"status": "success"}, - offset_seconds=3, - ), - _event( - "call-2", - event_type=EventType.TOOL_CALL, - metadata={"call_id": "two"}, - input_payload={"format": "csv"}, - offset_seconds=2, - ), - ] - - assert ToolExtractor.calculate_tool_success_rate("BACKUP", events) == 0.5 - candidate = ToolExtractor.extract_tool_pattern(events)[0] - assert candidate.metadata["total_count"] == 2 - assert candidate.metadata["unknown_count"] == 0 - assert candidate.metadata["duration_sample_count"] == 2 - - -def test_tool_success_rate_preserves_unknown_outcomes() -> None: - events = [ - _event("call-only", event_type=EventType.TOOL_CALL, input_payload={"mode": "dry-run"}), - ] - - candidate = ToolExtractor.extract_tool_pattern(events)[0] - assert ToolExtractor.calculate_tool_success_rate("backup", events) == 0.0 - assert candidate.metadata["total_count"] == 0 - assert candidate.metadata["unknown_count"] == 1 - assert candidate.metadata["data_completeness"] == 0.0 - - -def test_tool_pattern_infers_terminal_status_and_elapsed_duration() -> None: - events = [ - _event( - "sync-call", - event_type=EventType.TOOL_CALL, - tool_name="sync", - input_payload={"options": {"retry": True}, "secret": "must-not-leak"}, - offset_seconds=10, - ), - _event( - "sync-result", - event_type=EventType.TOOL_RESULT, - tool_name="sync", - output_payload={"status": "completed"}, - offset_seconds=12, - ), - _event( - "sync-orphan-failure", - event_type=EventType.TOOL_RESULT, - tool_name="sync", - output_payload={"status": "error", "error": "permission denied"}, - offset_seconds=14, - ), - ] - - candidate = ToolExtractor.extract_tool_pattern(events)[0] - assert candidate.metadata["success_count"] == 1 - assert candidate.metadata["failure_count"] == 1 - assert candidate.metadata["success_rate"] == 0.5 - assert candidate.metadata["average_response_ms"] == 2000.0 - assert candidate.metadata["common_failure_reasons"][0]["reason"].startswith("error: permission denied") - assert candidate.metadata["common_parameter_combinations"][0]["parameters"] == {"options.retry": "True"} +def test_extract_tool_pattern_uses_llm(tool_client: ToolFakeLLMClient) -> None: + events = [_event("call-1", EventType.TOOL_CALL, "backup"), _event("result-1", EventType.TOOL_RESULT, "backup", True)] + + candidates = ToolExtractor.extract_tool_pattern(events) + + assert tool_client.calls[0]["request"]["mode"] == "tool_pattern" + assert candidates[0].memory_type is MemoryType.TOOL + assert candidates[0].key == "tool.success_rate.backup" + + +def test_calculate_tool_success_rate_reads_llm_metadata(tool_client: ToolFakeLLMClient) -> None: + class RateClient(ToolFakeLLMClient): + def complete_json(self, prompt: str, schema: dict[str, Any]) -> dict[str, Any]: + payload = super().complete_json(prompt, schema) + payload["candidates"][0]["content"] = "backup success rate" + payload["candidates"][0]["confidence"] = 0.9 + payload["candidates"][0]["metadata"] = {"tool_name": "backup", "success_rate": 2 / 3} + return payload + + client = RateClient() + set_default_llm_client(client) + try: + rate = ToolExtractor.calculate_tool_success_rate("backup", [_event("result-1", EventType.TOOL_RESULT, "backup", True)]) + finally: + clear_default_llm_client() + + assert client.calls[0]["request"]["mode"] == "tool_success_rate:backup" + assert rate == pytest.approx(2 / 3) + + +def test_tool_extractor_has_no_rule_fallback_without_client() -> None: + clear_default_llm_client() + + assert ToolExtractor.extract_tool_pattern([_event("result-1", EventType.TOOL_RESULT, "backup", True)]) == [] + assert ToolExtractor.calculate_tool_success_rate("backup", []) == 0.0 + + +def test_calculate_tool_success_rate_returns_zero_when_llm_metadata_is_missing_or_mismatched() -> None: + class MissingRateClient(ToolFakeLLMClient): + def complete_json(self, prompt: str, schema: dict[str, Any]) -> dict[str, Any]: + payload = super().complete_json(prompt, schema) + payload["candidates"][0]["metadata"] = {"tool_name": "other_tool"} + return payload + + set_default_llm_client(MissingRateClient()) + try: + assert ToolExtractor.calculate_tool_success_rate("backup", [_event("result-1", EventType.TOOL_RESULT, "backup", True)]) == 0.0 + finally: + clear_default_llm_client() diff --git a/tests/test_workflow_extractor.py b/tests/test_workflow_extractor.py index be777f0..74d9397 100644 --- a/tests/test_workflow_extractor.py +++ b/tests/test_workflow_extractor.py @@ -1,482 +1,100 @@ from __future__ import annotations +import json from datetime import datetime, timezone +from typing import Any + +import pytest from core.constants import EventType, MemoryType, Scene from core.models import MemoryEvent -from extractors import workflow_extractor as workflow_mod +from extractors.llm_memory_extractor import clear_default_llm_client, set_default_llm_client from extractors.workflow_extractor import WorkflowExtractor -def _event( - event_id: str, - *, - user_id: str = "user-1", - session_id: str = "session-1", - task_id: str = "task-1", - event_type: EventType = EventType.CONVERSATION, - scenario: Scene = Scene.GLOBAL, - source: str = "conversation", - content: str | None = None, - tool_name: str | None = None, - input_payload: dict | None = None, - output_payload: dict | None = None, - metadata: dict | None = None, - success: bool | None = True, - timestamp: datetime | None = None, -) -> MemoryEvent: +class WorkflowFakeLLMClient: + def __init__(self) -> None: + self.calls: list[dict[str, Any]] = [] + + def complete_json(self, prompt: str, schema: dict[str, Any]) -> dict[str, Any]: + request = json.loads(prompt) + self.calls.append({"request": request, "schema": schema}) + if "boundaries" in schema.get("required", []): + return {"boundaries": [{"start": 0, "end": 2, "reason": "同一工作流"}]} + return { + "candidates": [ + { + "is_memory_worthy": True, + "is_long_term": True, + "memory_type": "workflow", + "category": "monthly_report", + "value": "read_generate_export", + "scope": "office", + "content": "月报流程:读取数据 -> 生成报告 -> 导出 PDF", + "confidence": 0.86, + "evidence": "LLM reviewed the ordered task events.", + "reason": "多个事件共同构成可复用流程。", + "sensitivity": "none", + } + ] + } + + +@pytest.fixture() +def workflow_client() -> WorkflowFakeLLMClient: + client = WorkflowFakeLLMClient() + set_default_llm_client(client) + try: + yield client + finally: + clear_default_llm_client() + + +def _event(event_id: str, event_type: EventType = EventType.CONVERSATION, tool_name: str | None = None) -> MemoryEvent: return MemoryEvent( event_id=event_id, raw_event_id=f"raw-{event_id}", - user_id=user_id, - session_id=session_id, - task_id=task_id, + user_id="u-flow", + session_id="s-flow", + task_id="t-flow", event_type=event_type, - scenario=scenario, - source=source, - content=content, + scenario=Scene.OFFICE, + source=event_type.value, + actor="agent", + content="处理月报" if event_type is EventType.CONVERSATION else None, tool_name=tool_name, - input=input_payload or {}, - output=output_payload or {}, - metadata=metadata or {}, - success=success, - timestamp=timestamp or datetime(2026, 6, 23, tzinfo=timezone.utc), - raw_event={"event_id": event_id}, + timestamp=datetime(2026, 7, 5, 10, 0, tzinfo=timezone.utc), ) -def test_detect_workflow_boundary_splits_separate_clusters() -> None: - events = [ - _event("evt-0", content="just chatting"), - _event( - "evt-1", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="download", - input_payload={"url": "https://example.com/a.csv"}, - output_payload={"file": "a.csv"}, - ), - _event( - "evt-2", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="download", - output_payload={"file": "a.csv"}, - ), - _event("evt-3", content="small talk"), - _event("evt-4", content="still small talk"), - _event( - "evt-5", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="export", - input_payload={"source": "a.csv"}, - output_payload={"file": "report.xlsx"}, - ), - _event( - "evt-6", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="export", - output_payload={"file": "report.xlsx"}, - ), - ] +def test_detect_workflow_boundary_uses_llm(workflow_client: WorkflowFakeLLMClient) -> None: + events = [_event("e0"), _event("e1", EventType.TOOL_CALL, "read_sheet"), _event("e2", EventType.TOOL_RESULT, "export_pdf")] - assert WorkflowExtractor.detect_workflow_boundary(events) == [(1, 2), (5, 6)] + assert WorkflowExtractor.detect_workflow_boundary(events) == [(0, 2)] + assert workflow_client.calls[0]["request"]["mode"] == "workflow_boundary_detection" -def test_extract_tool_sequence_recognizes_continuous_tool_chain() -> None: - events = [ - _event("evt-0", content="please handle this in order"), - _event( - "evt-1", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="download", - input_payload={"url": "https://example.com/raw.csv"}, - output_payload={"file": "raw.csv"}, - ), - _event( - "evt-2", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="download", - output_payload={"file": "raw.csv"}, - ), - _event("evt-3", content="use raw.csv next"), - _event( - "evt-4", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="transform", - input_payload={"source": "raw.csv"}, - output_payload={"file": "processed.csv"}, - ), - _event( - "evt-5", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="transform", - output_payload={"file": "processed.csv"}, - ), - _event("evt-6", content="then export the output"), - _event( - "evt-7", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="export", - input_payload={"input": "processed.csv"}, - output_payload={"file": "report.xlsx"}, - ), - _event( - "evt-8", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="export", - output_payload={"file": "report.xlsx"}, - ), - ] +def test_extract_tool_sequence_uses_llm(workflow_client: WorkflowFakeLLMClient) -> None: + events = [_event("e0"), _event("e1", EventType.TOOL_CALL, "read_sheet"), _event("e2", EventType.TOOL_RESULT, "export_pdf")] candidates = WorkflowExtractor.extract_tool_sequence(events) - assert len(candidates) == 1 - candidate = candidates[0] - assert candidate.memory_type is MemoryType.WORKFLOW - assert candidate.scenario is Scene.GLOBAL - assert candidate.key.startswith("workflow.tool_sequence.") - assert candidate.content.startswith("Tool sequence") - assert candidate.metadata["pattern"] == "tool_sequence" - assert candidate.metadata["step_count"] == 3 - assert candidate.metadata["tool_names"] == ["download", "transform", "export"] - assert candidate.metadata["reproduction_rate"] >= 0.8 - assert len(candidate.metadata["dependencies"]) == 2 - assert candidate.metadata["dependencies"][0]["from_tool"] == "download" - assert candidate.metadata["dependencies"][0]["to_tool"] == "transform" - assert candidate.metadata["dependencies"][1]["from_tool"] == "transform" - assert candidate.metadata["dependencies"][1]["to_tool"] == "export" - - -def test_extract_multi_step_workflow_recognizes_dependencies() -> None: - events = [ - _event("evt-0", content="step by step, use the previous result"), - _event( - "evt-1", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="search", - input_payload={"query": "dataset"}, - output_payload={"file": "dataset.csv"}, - ), - _event( - "evt-2", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="search", - output_payload={"file": "dataset.csv"}, - ), - _event("evt-3", content="then use dataset.csv for the next step"), - _event( - "evt-4", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="analyse", - input_payload={"input": "dataset.csv"}, - output_payload={"file": "analysis.json"}, - ), - _event( - "evt-5", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="analyse", - output_payload={"file": "analysis.json"}, - ), - _event("evt-6", content="after that, export the report"), - _event( - "evt-7", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="export", - input_payload={"input": "analysis.json"}, - output_payload={"file": "report.xlsx"}, - ), - _event( - "evt-8", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="export", - output_payload={"file": "report.xlsx"}, - ), - ] - - candidates = WorkflowExtractor.extract_multi_step_workflow(events) - - assert len(candidates) == 1 - candidate = candidates[0] - assert candidate.memory_type is MemoryType.WORKFLOW - assert candidate.scenario is Scene.GLOBAL - assert candidate.key.startswith("workflow.multi_step.") - assert candidate.content.startswith("Complex workflow") - assert candidate.metadata["pattern"] == "multi_step" - assert candidate.metadata["step_count"] == 3 - assert candidate.metadata["reproduction_rate"] >= 0.8 - assert len(candidate.metadata["dependencies"]) == 2 - assert candidate.metadata["dependencies"][0]["evidence"] - assert candidate.metadata["dependencies"][1]["evidence"] - - -def test_extract_tool_sequence_deduplicates_repeated_workflows() -> None: - workflow_one = [ - _event( - "evt-a1", - session_id="session-a", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="download", - input_payload={"url": "https://example.com/raw.csv"}, - output_payload={"file": "raw.csv"}, - ), - _event( - "evt-a2", - session_id="session-a", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="download", - output_payload={"file": "raw.csv"}, - ), - _event("evt-a3", session_id="session-a", content="use raw.csv next"), - _event( - "evt-a4", - session_id="session-a", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="transform", - input_payload={"source": "raw.csv"}, - output_payload={"file": "processed.csv"}, - ), - _event( - "evt-a5", - session_id="session-a", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="transform", - output_payload={"file": "processed.csv"}, - ), - ] - workflow_two = [ - _event( - "evt-b1", - session_id="session-b", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="download", - input_payload={"url": "https://example.com/raw.csv"}, - output_payload={"file": "raw.csv"}, - ), - _event( - "evt-b2", - session_id="session-b", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="download", - output_payload={"file": "raw.csv"}, - ), - _event("evt-b3", session_id="session-b", content="use raw.csv next"), - _event( - "evt-b4", - session_id="session-b", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="transform", - input_payload={"source": "raw.csv"}, - output_payload={"file": "processed.csv"}, - ), - _event( - "evt-b5", - session_id="session-b", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="transform", - output_payload={"file": "processed.csv"}, - ), - ] - - candidates = WorkflowExtractor.extract_tool_sequence(workflow_one + workflow_two) - - assert len(candidates) == 1 - assert candidates[0].metadata["workflow_signature"] == "download__transform" - assert candidates[0].metadata["occurrence_count"] == 2 - assert set(candidates[0].source_events) == {event.event_id for event in workflow_one + workflow_two} - + assert workflow_client.calls[0]["request"]["mode"] == "workflow_tool_sequence" + assert candidates[0].memory_type is MemoryType.WORKFLOW + assert candidates[0].key == "workflow.monthly_report.read_generate_export" -def test_workflow_extraction_orders_events_and_isolates_concurrent_tasks() -> None: - base = datetime(2026, 6, 23, tzinfo=timezone.utc) - task_one = [ - _event( - "one-download-call", - task_id="task-one", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="download", - output_payload={"file": "one.csv"}, - timestamp=base.replace(second=1), - ), - _event( - "one-download-result", - task_id="task-one", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="download", - output_payload={"file": "one.csv"}, - timestamp=base.replace(second=2), - ), - _event( - "one-export-call", - task_id="task-one", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="export", - input_payload={"source": "one.csv"}, - output_payload={"file": "one.xlsx"}, - timestamp=base.replace(second=3), - ), - _event( - "one-export-result", - task_id="task-one", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="export", - output_payload={"file": "one.xlsx"}, - timestamp=base.replace(second=4), - ), - ] - task_two = [ - _event( - "two-search-call", - task_id="task-two", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="search", - output_payload={"file": "two.csv"}, - timestamp=base.replace(second=1), - ), - _event( - "two-search-result", - task_id="task-two", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="search", - output_payload={"file": "two.csv"}, - timestamp=base.replace(second=2), - ), - _event( - "two-analyse-call", - task_id="task-two", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="analyse", - input_payload={"source": "two.csv"}, - output_payload={"file": "two.json"}, - timestamp=base.replace(second=3), - ), - _event( - "two-analyse-result", - task_id="task-two", - event_type=EventType.TOOL_RESULT, - source="tool", - tool_name="analyse", - output_payload={"file": "two.json"}, - timestamp=base.replace(second=4), - ), - ] - candidates = WorkflowExtractor.extract_tool_sequence(list(reversed(task_one + task_two))) +def test_extract_multi_step_workflow_uses_llm(workflow_client: WorkflowFakeLLMClient) -> None: + candidates = WorkflowExtractor.extract_multi_step_workflow([_event("e0"), _event("e1", EventType.TOOL_CALL, "read_sheet")]) - assert len(candidates) == 2 - by_steps = {tuple(candidate.metadata["tool_names"]): candidate for candidate in candidates} - assert set(by_steps) == {("download", "export"), ("search", "analyse")} - assert set(by_steps[("download", "export")].source_events) == {event.event_id for event in task_one} - assert set(by_steps[("search", "analyse")].source_events) == {event.event_id for event in task_two} - assert all(candidate.metadata["reproduction_rate"] >= 0.8 for candidate in candidates) + assert workflow_client.calls[0]["request"]["mode"] == "workflow_multi_step" + assert candidates[0].source == "llm_extracted" -def test_workflow_requires_real_reconstruction_evidence() -> None: - events = [ - _event("call-1", event_type=EventType.TOOL_CALL, source="tool", tool_name="first"), - _event("result-1", event_type=EventType.TOOL_RESULT, source="tool", tool_name="first"), - _event("call-2", event_type=EventType.TOOL_CALL, source="tool", tool_name="second"), - _event("result-2", event_type=EventType.TOOL_RESULT, source="tool", tool_name="second"), - ] +def test_workflow_extractor_has_no_rule_fallback_without_client() -> None: + clear_default_llm_client() + events = [_event("e0"), _event("e1", EventType.TOOL_CALL, "read_sheet")] + assert WorkflowExtractor.detect_workflow_boundary(events) == [] assert WorkflowExtractor.extract_tool_sequence(events) == [] - - -def test_workflow_internal_helper_branches_and_empty_paths() -> None: - assert workflow_mod._iter_text_fragments(None) == [] - assert workflow_mod._iter_text_fragments(" abc ") == ["abc"] - assert workflow_mod._iter_text_fragments(5) == ["5"] - - fragments = workflow_mod._iter_text_fragments( - { - "skip": "value", - "nested": {"x": "y"}, - "items": [1, 2], - "flags": {True, False}, - } - ) - assert "y" in fragments - assert "1" in fragments and "2" in fragments - assert "True" in fragments and "False" in fragments - - relevant = _event("evt-relevant", content="step 1 then next") - irrelevant = _event("evt-irrelevant", content="plain chat") - assert workflow_mod._is_workflow_relevant(relevant) - assert not workflow_mod._is_workflow_relevant(irrelevant) - - tool_call = _event("evt-tool-call", event_type=EventType.TOOL_CALL, source="tool", tool_name="Bash") - tool_result = _event("evt-tool-result", event_type=EventType.TOOL_RESULT, source="tool", tool_name="Bash") - tool_other = _event("evt-tool-other", event_type=EventType.TOOL_CALL, source="tool", tool_name="Cat") - assert workflow_mod._same_tool_group(tool_call, tool_result) - assert not workflow_mod._same_tool_group(tool_call, tool_other) - - assert workflow_mod._step_label([tool_call]) == "bash" - assert workflow_mod._step_label([_event("evt-label", content=" Use the output ")]) == "use the output" - assert workflow_mod._step_label([_event("evt-fallback", content=None, tool_name=None)]).startswith("conversation") - - assert workflow_mod._group_transition_markers([_event("evt-trans", content="then use it")]) - assert not workflow_mod._group_transition_markers([_event("evt-notrans", content="just text")]) - - simple_events = [ - _event("evt-b0", event_type=EventType.TOOL_CALL, source="tool", tool_name="download", output_payload={"file": "a.csv"}), - _event("evt-b1", event_type=EventType.TOOL_RESULT, source="tool", tool_name="download", output_payload={"file": "a.csv"}), - _event("evt-b2", content="after that"), - _event("evt-b3", event_type=EventType.TOOL_CALL, source="tool", tool_name="export", input_payload={"source": "a.csv"}), - ] - assert workflow_mod._boundary_signature(simple_events, [(0, 1), (3, 3)]).startswith("download") - assert workflow_mod._workflow_content("Prefix", [[tool_call, tool_result], [tool_other]], []) == "Prefix: 1. bash | 2. cat" - assert "deps:" in workflow_mod._workflow_content( - "Prefix", - [[tool_call], [tool_other]], - [{"from_tool": "bash", "to_tool": "cat"}], - ) - assert workflow_mod._workflow_confidence(2, 0, False) < workflow_mod._workflow_confidence(4, 2, True) - - dependency_left = _event( - "evt-left", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="download", - output_payload={"file": "a.csv"}, - ) - dependency_right = _event( - "evt-right", - event_type=EventType.TOOL_CALL, - source="tool", - tool_name="transform", - input_payload={"source": "a.csv"}, - ) - edges = workflow_mod._dependency_edges([[dependency_left], [dependency_right]]) - assert edges and edges[0]["from_tool"] == "download" and edges[0]["to_tool"] == "transform" - - assert WorkflowExtractor.detect_workflow_boundary([]) == [] - assert WorkflowExtractor.extract_tool_sequence([tool_call, tool_result]) == [] - assert WorkflowExtractor.extract_multi_step_workflow([tool_call, tool_result]) == [] + assert WorkflowExtractor.extract_multi_step_workflow(events) == []