text2mem_practice.py 6.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159
  1. """
  2. Text2Mem 实战:实现一个最小可用的记忆操作引擎
  3. 运行:python text2mem_practice.py
  4. """
  5. import json
  6. from datetime import datetime
  7. from typing import ClassVar, Literal
  8. from pydantic import BaseModel, Field, field_validator, model_validator
  9. # ===== 结构与业务校验(生产环境可额外发布 JSON Schema) =====
  10. class MemMeta(BaseModel):
  11. timestamp: str = ""
  12. source: str = "user_direct"
  13. confidence: float = 1.0
  14. dry_run: bool = False
  15. confirmation: bool = False
  16. def model_post_init(self, __context):
  17. if not self.timestamp:
  18. self.timestamp = datetime.now().isoformat()
  19. class MemOp(BaseModel):
  20. """五元组 IR 操作"""
  21. stage: Literal["ENC", "RET", "STO"]
  22. op: str
  23. target: str
  24. args: dict = Field(default_factory=dict)
  25. meta: MemMeta = Field(default_factory=MemMeta)
  26. VALID_OPS: ClassVar[dict[str, set[str]]] = {
  27. "ENC": {"Encode"},
  28. "STO": {"Update", "Label", "Promote", "Demote", "Merge", "Split", "Delete", "Lock", "Expire"},
  29. "RET": {"Retrieve", "Summarize"},
  30. }
  31. @field_validator("op")
  32. @classmethod
  33. def op_must_be_valid(cls, v):
  34. all_ops = set()
  35. for ops in cls.VALID_OPS.values():
  36. all_ops |= ops
  37. if v not in all_ops:
  38. raise ValueError(f"Unknown op: {v}. Valid: {sorted(all_ops)}")
  39. return v
  40. @model_validator(mode="after")
  41. def op_stage_match(self):
  42. valid = self.VALID_OPS.get(self.stage, set())
  43. if self.op not in valid:
  44. raise ValueError(f"Op '{self.op}' not allowed in stage '{self.stage}'. Valid: {sorted(valid)}")
  45. return self
  46. # ===== 第二层:业务逻辑校验 =====
  47. class MemoryStore:
  48. """模拟记忆存储引擎;仅实现 Encode/Update/Lock/Delete/Retrieve。"""
  49. WRITE_OPS: ClassVar[set[str]] = {"Encode", "Update", "Lock", "Delete"}
  50. def __init__(self):
  51. self.memories: dict[str, dict] = {}
  52. def execute(self, ir: MemOp) -> dict:
  53. if ir.meta.dry_run:
  54. return {"status": "dry_run", "would_execute": ir.model_dump()}
  55. if ir.op in self.WRITE_OPS and not ir.meta.confirmation:
  56. return {"status": "pending_confirmation", "msg": "Write requires external confirmation"}
  57. handler = getattr(self, f"_exec_{ir.op.lower()}", None)
  58. if not handler:
  59. return {"status": "error", "msg": f"No handler for op '{ir.op}'"}
  60. return handler(ir)
  61. def _exec_encode(self, ir: MemOp) -> dict:
  62. mem_id = ir.target or f"mem_{len(self.memories) + 1:04d}"
  63. self.memories[mem_id] = {
  64. "id": mem_id, "content": ir.args.get("content", ""),
  65. "priority": ir.args.get("priority", "normal"),
  66. "locked": False, "tags": [],
  67. "created_at": ir.meta.timestamp, "updated_at": ir.meta.timestamp,
  68. }
  69. return {"status": "ok", "mem_id": mem_id}
  70. def _exec_update(self, ir: MemOp) -> dict:
  71. mem = self._get_or_error(ir.target)
  72. if isinstance(mem, dict) and "error" in mem:
  73. return mem
  74. if mem["locked"]:
  75. return {"status": "error", "msg": f"Memory '{ir.target}' is locked, cannot update"}
  76. fields = ir.args.get("fields", {})
  77. mem.update(fields)
  78. mem["updated_at"] = ir.meta.timestamp
  79. return {"status": "ok", "updated": list(fields.keys())}
  80. def _exec_lock(self, ir: MemOp) -> dict:
  81. mem = self._get_or_error(ir.target)
  82. if isinstance(mem, dict) and "error" in mem:
  83. return mem
  84. mem["locked"] = True
  85. return {"status": "ok", "locked": ir.target}
  86. def _exec_delete(self, ir: MemOp) -> dict:
  87. mem = self._get_or_error(ir.target)
  88. if isinstance(mem, dict) and "error" in mem:
  89. return mem
  90. if mem["locked"]:
  91. return {"status": "error", "msg": f"Memory '{ir.target}' is locked, cannot delete"}
  92. del self.memories[ir.target]
  93. return {"status": "ok", "deleted": ir.target}
  94. def _exec_retrieve(self, ir: MemOp) -> dict:
  95. keyword = ir.args.get("keyword", "")
  96. results = [m for m in self.memories.values() if keyword.lower() in m["content"].lower()]
  97. return {"status": "ok", "count": len(results), "results": results}
  98. def _get_or_error(self, target: str):
  99. if target not in self.memories:
  100. return {"status": "error", "msg": f"Memory '{target}' not found"}
  101. return self.memories[target]
  102. # ===== 运行演示 =====
  103. if __name__ == "__main__":
  104. store = MemoryStore()
  105. ops = [
  106. MemOp(stage="ENC", op="Encode", target="mem_user_lang",
  107. args={"content": "用户偏好 Python 写后端,TypeScript 写前端", "priority": "high"},
  108. meta=MemMeta(confirmation=True)),
  109. MemOp(stage="ENC", op="Encode", target="mem_project",
  110. args={"content": "当前项目:企业内部智能知识库,使用通义千问和 RAG"},
  111. meta=MemMeta(confirmation=True)),
  112. MemOp(stage="STO", op="Lock", target="mem_user_lang",
  113. args={"reason": "用户身份信息,禁止自动修改"},
  114. meta=MemMeta(confirmation=True)),
  115. MemOp(stage="STO", op="Update", target="mem_user_lang",
  116. args={"fields": {"content": "被篡改的内容"}}),
  117. MemOp(stage="STO", op="Update", target="mem_user_lang",
  118. args={"fields": {"content": "仍应被锁拦截"}},
  119. meta=MemMeta(confirmation=True)),
  120. MemOp(stage="STO", op="Update", target="mem_project",
  121. args={"fields": {"content": "当前项目:企业内部智能知识库,检索重排进度 75%"}},
  122. meta=MemMeta(confirmation=True)),
  123. MemOp(stage="RET", op="Retrieve", target="", args={"keyword": "知识库"}),
  124. MemOp(stage="STO", op="Delete", target="mem_project",
  125. args={}, meta=MemMeta(dry_run=True)),
  126. ]
  127. print("=" * 60)
  128. print("Text2Mem 实战:记忆操作引擎演示")
  129. print("=" * 60)
  130. for i, op in enumerate(ops, 1):
  131. print(f"\n--- 操作 {i}: {op.stage}/{op.op} ---")
  132. result = store.execute(op)
  133. print(f" 结果: {json.dumps(result, ensure_ascii=False, indent=2)}")
  134. print(f"\n最终记忆库状态({len(store.memories)} 条):")
  135. for mid, mem in store.memories.items():
  136. locked = " 🔒" if mem["locked"] else ""
  137. print(f" [{mid}]{locked}: {mem['content']}")