pycontextdb 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (47) hide show
  1. contextdb/__init__.py +72 -0
  2. contextdb/agents/__init__.py +8 -0
  3. contextdb/agents/memory_bus.py +80 -0
  4. contextdb/agents/rl_manager.py +89 -0
  5. contextdb/cli.py +107 -0
  6. contextdb/client.py +516 -0
  7. contextdb/core/__init__.py +41 -0
  8. contextdb/core/config.py +89 -0
  9. contextdb/core/exceptions.py +29 -0
  10. contextdb/core/models.py +151 -0
  11. contextdb/dynamics/__init__.py +25 -0
  12. contextdb/dynamics/evolution.py +168 -0
  13. contextdb/dynamics/formation.py +193 -0
  14. contextdb/dynamics/retrieval.py +130 -0
  15. contextdb/graphs/__init__.py +17 -0
  16. contextdb/graphs/base.py +46 -0
  17. contextdb/graphs/causal.py +224 -0
  18. contextdb/graphs/entity.py +251 -0
  19. contextdb/graphs/semantic.py +156 -0
  20. contextdb/graphs/temporal.py +173 -0
  21. contextdb/integrations/__init__.py +10 -0
  22. contextdb/integrations/autogen.py +39 -0
  23. contextdb/integrations/crewai.py +41 -0
  24. contextdb/integrations/langchain.py +132 -0
  25. contextdb/integrations/openai_tools.py +124 -0
  26. contextdb/memory/__init__.py +9 -0
  27. contextdb/memory/experiential.py +102 -0
  28. contextdb/memory/factual.py +58 -0
  29. contextdb/memory/working.py +90 -0
  30. contextdb/privacy/__init__.py +9 -0
  31. contextdb/privacy/audit.py +199 -0
  32. contextdb/privacy/pii_detector.py +173 -0
  33. contextdb/privacy/retention.py +99 -0
  34. contextdb/py.typed +0 -0
  35. contextdb/store/__init__.py +15 -0
  36. contextdb/store/base.py +67 -0
  37. contextdb/store/sqlite_store.py +517 -0
  38. contextdb/store/vector_index.py +241 -0
  39. contextdb/utils/__init__.py +22 -0
  40. contextdb/utils/embeddings.py +159 -0
  41. contextdb/utils/llm.py +139 -0
  42. contextdb/utils/migrations.py +159 -0
  43. pycontextdb-0.1.0.dist-info/METADATA +589 -0
  44. pycontextdb-0.1.0.dist-info/RECORD +47 -0
  45. pycontextdb-0.1.0.dist-info/WHEEL +4 -0
  46. pycontextdb-0.1.0.dist-info/entry_points.txt +2 -0
  47. pycontextdb-0.1.0.dist-info/licenses/LICENSE +190 -0
@@ -0,0 +1,151 @@
1
+ """Core data models for ContextDB.
2
+
3
+ Every persisted object in ContextDB is modeled here. These Pydantic v2 types
4
+ are the canonical representation used across storage, graphs, and the public
5
+ API — so please keep them backwards-compatible when amending.
6
+
7
+ Design notes:
8
+
9
+ * All timestamps are timezone-aware UTC (``datetime.now(tz=timezone.utc)``).
10
+ Storing naive times is a recipe for silent drift between nodes.
11
+ * Enums inherit from ``str`` so that ``model_dump()`` and SQL round-trips
12
+ produce human-readable values.
13
+ * Default-factory is used for all mutable defaults (dict / list) — never bare
14
+ ``= {}`` or ``= []`` — so shared-state bugs cannot appear.
15
+ """
16
+
17
+ from __future__ import annotations
18
+
19
+ from datetime import datetime, timedelta, timezone
20
+ from enum import Enum
21
+ from typing import Any, Literal
22
+ from uuid import uuid4
23
+
24
+ from pydantic import BaseModel, Field
25
+
26
+ GraphType = Literal["semantic", "temporal", "causal", "entity"]
27
+
28
+
29
+ def _utcnow() -> datetime:
30
+ """Return the current UTC time as a timezone-aware :class:`datetime`."""
31
+ return datetime.now(tz=timezone.utc)
32
+
33
+
34
+ class MemoryType(str, Enum):
35
+ """High-level category a :class:`MemoryItem` belongs to."""
36
+
37
+ FACTUAL = "FACTUAL"
38
+ EXPERIENTIAL = "EXPERIENTIAL"
39
+ WORKING = "WORKING"
40
+
41
+
42
+ class MemoryStatus(str, Enum):
43
+ """Lifecycle state of a memory item."""
44
+
45
+ ACTIVE = "ACTIVE"
46
+ ARCHIVED = "ARCHIVED"
47
+ DELETED = "DELETED"
48
+
49
+
50
+ class PIIType(str, Enum):
51
+ """Recognized PII categories. ``CUSTOM`` is an escape hatch for users."""
52
+
53
+ NAME = "NAME"
54
+ EMAIL = "EMAIL"
55
+ PHONE = "PHONE"
56
+ ADDRESS = "ADDRESS"
57
+ SSN = "SSN"
58
+ CREDIT_CARD = "CREDIT_CARD"
59
+ CUSTOM = "CUSTOM"
60
+
61
+
62
+ class PIIAnnotation(BaseModel):
63
+ """A single PII span detected within ``MemoryItem.content``.
64
+
65
+ ``start`` and ``end`` are character offsets (half-open, Python slice
66
+ semantics) into the **original** content — not the redacted form.
67
+ """
68
+
69
+ pii_type: PIIType
70
+ start: int = Field(ge=0, description="Character offset (inclusive).")
71
+ end: int = Field(ge=0, description="Character offset (exclusive).")
72
+ original: str = Field(description="Original text that was flagged.")
73
+ redacted: str = Field(description="Replacement text (e.g., '[NAME]').")
74
+
75
+
76
+ class Edge(BaseModel):
77
+ """A directed edge between two memory items in one of four graphs."""
78
+
79
+ source_id: str
80
+ target_id: str
81
+ graph_type: GraphType
82
+ weight: float = 1.0
83
+ metadata: dict[str, Any] = Field(default_factory=dict)
84
+ created_at: datetime = Field(default_factory=_utcnow)
85
+
86
+
87
+ class Entity(BaseModel):
88
+ """A named entity extracted from one or more memories."""
89
+
90
+ name: str
91
+ entity_type: str = Field(description="e.g., PERSON, ORG, PRODUCT, LOCATION.")
92
+ attributes: dict[str, Any] = Field(default_factory=dict)
93
+ memory_ids: list[str] = Field(
94
+ default_factory=list,
95
+ description="IDs of memories that mention this entity.",
96
+ )
97
+
98
+
99
+ class RetentionPolicy(BaseModel):
100
+ """Declarative retention rules applied by the retention enforcer.
101
+
102
+ Each ``*_ttl`` field may be ``None`` to disable expiry for that class of
103
+ memory. The default policy is chosen to match common privacy expectations:
104
+ long for factual, unbounded for experiential, short for working.
105
+ """
106
+
107
+ default_ttl: timedelta | None = timedelta(days=730)
108
+ factual_ttl: timedelta | None = timedelta(days=1825)
109
+ experiential_ttl: timedelta | None = None
110
+ working_ttl: timedelta | None = timedelta(hours=24)
111
+ right_to_erasure: bool = True
112
+
113
+
114
+ class MemoryItem(BaseModel):
115
+ """The canonical unit of memory in ContextDB.
116
+
117
+ A :class:`MemoryItem` carries its content, vector embedding, lifecycle
118
+ metadata, privacy annotations, and back-references to entities and tags.
119
+ Graph relationships live in :class:`Edge` objects, not on the item
120
+ directly, so a memory can participate in multiple graphs without
121
+ schema churn.
122
+ """
123
+
124
+ id: str = Field(default_factory=lambda: str(uuid4()))
125
+ content: str
126
+ embedding: list[float] | None = None
127
+ memory_type: MemoryType = MemoryType.FACTUAL
128
+ source: str = ""
129
+ metadata: dict[str, Any] = Field(default_factory=dict)
130
+
131
+ event_time: datetime | None = Field(
132
+ default=None,
133
+ description="When the event occurred (valid-time). Distinct from ingestion_time.",
134
+ )
135
+ ingestion_time: datetime = Field(
136
+ default_factory=_utcnow,
137
+ description="When ContextDB stored the memory (system-time).",
138
+ )
139
+
140
+ pii_annotations: list[PIIAnnotation] = Field(default_factory=list)
141
+ retention_policy: RetentionPolicy | None = None
142
+
143
+ created_at: datetime = Field(default_factory=_utcnow)
144
+ updated_at: datetime = Field(default_factory=_utcnow)
145
+ access_count: int = 0
146
+ last_accessed: datetime | None = None
147
+ confidence: float = 1.0
148
+ status: MemoryStatus = MemoryStatus.ACTIVE
149
+
150
+ entity_mentions: list[str] = Field(default_factory=list)
151
+ tags: list[str] = Field(default_factory=list)
@@ -0,0 +1,25 @@
1
+ """Dynamics layer — formation, evolution, retrieval of memories over time."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from contextdb.dynamics.evolution import AutoLinker, Consolidator, Pruner
6
+ from contextdb.dynamics.formation import (
7
+ FormationPipeline,
8
+ MemoryCompressor,
9
+ MemoryExtractor,
10
+ Segmenter,
11
+ )
12
+ from contextdb.dynamics.retrieval import QueryClassifier, RetrievalEngine, RetrievalFuser
13
+
14
+ __all__ = [
15
+ "AutoLinker",
16
+ "Consolidator",
17
+ "FormationPipeline",
18
+ "MemoryCompressor",
19
+ "MemoryExtractor",
20
+ "Pruner",
21
+ "QueryClassifier",
22
+ "RetrievalEngine",
23
+ "RetrievalFuser",
24
+ "Segmenter",
25
+ ]
@@ -0,0 +1,168 @@
1
+ """Evolution engine — auto-linking, consolidation, and pruning.
2
+
3
+ * :class:`AutoLinker` mirrors a new memory into every configured graph's
4
+ ``add_node`` hook so edges appear as memories are written.
5
+ * :class:`Consolidator` finds dense semantic clusters and replaces them with
6
+ a single summary memory, pruning the originals (archive, not delete).
7
+ * :class:`Pruner` applies decay / age / redundancy strategies to drop stale
8
+ memories — the operator chooses the strategy per call.
9
+
10
+ The consolidator is intentionally conservative: it only touches memories
11
+ that have multiple neighbors above the semantic threshold, so a single
12
+ outlier never triggers a merge.
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ from datetime import datetime, timedelta, timezone
18
+ from typing import TYPE_CHECKING, Any
19
+
20
+ from contextdb.core.models import MemoryItem, MemoryStatus, MemoryType
21
+
22
+ if TYPE_CHECKING:
23
+ from contextdb.graphs.base import BaseGraph
24
+ from contextdb.graphs.semantic import SemanticGraph
25
+ from contextdb.store.sqlite_store import SQLiteStore
26
+ from contextdb.utils.llm import LLMProvider
27
+
28
+
29
+ class AutoLinker:
30
+ """Forward newly-added nodes to each graph's ``add_node`` hook."""
31
+
32
+ def __init__(self, graphs: dict[str, BaseGraph]) -> None:
33
+ self.graphs = graphs
34
+
35
+ async def link(self, memory_id: str, data: dict[str, Any]) -> None:
36
+ for graph in self.graphs.values():
37
+ try:
38
+ await graph.add_node(memory_id, data)
39
+ except Exception: # noqa: BLE001
40
+ # One graph failing must not block the others. Auto-linking
41
+ # is best-effort; hard failures surface on explicit calls.
42
+ continue
43
+
44
+
45
+ class Consolidator:
46
+ """Merge dense semantic clusters into a single summary memory."""
47
+
48
+ def __init__(
49
+ self,
50
+ store: SQLiteStore,
51
+ semantic_graph: SemanticGraph,
52
+ llm: LLMProvider,
53
+ summary_prompt: str | None = None,
54
+ ) -> None:
55
+ self.store = store
56
+ self.semantic = semantic_graph
57
+ self.llm = llm
58
+ self.summary_prompt = summary_prompt or (
59
+ "Summarize these related memories into one coherent statement. "
60
+ "Preserve every entity and date.\n\n{memories}"
61
+ )
62
+
63
+ async def consolidate(self, min_cluster_size: int = 5) -> list[MemoryItem]:
64
+ """Walk active memories in 500-row pages, merging dense clusters."""
65
+ visited: set[str] = set()
66
+ summaries: list[MemoryItem] = []
67
+ async for memory in self.store.iter_memories(batch_size=500):
68
+ if memory.id in visited or memory.status != MemoryStatus.ACTIVE:
69
+ continue
70
+ neighbors = await self.semantic.get_neighbors(
71
+ memory.id, depth=1, max_results=min_cluster_size * 2
72
+ )
73
+ cluster_ids = [memory.id] + [nid for nid, _ in neighbors]
74
+ cluster_ids = [cid for cid in cluster_ids if cid not in visited]
75
+ if len(cluster_ids) < min_cluster_size:
76
+ continue
77
+ cluster_items: list[MemoryItem] = []
78
+ for cid in cluster_ids:
79
+ item = await self.store.get_raw(cid)
80
+ if item is None or item.status != MemoryStatus.ACTIVE:
81
+ continue
82
+ cluster_items.append(item)
83
+ if len(cluster_items) < min_cluster_size:
84
+ continue
85
+
86
+ summary = await self._summarize([m.content for m in cluster_items])
87
+ if not summary:
88
+ continue
89
+ new_item = MemoryItem(
90
+ content=summary,
91
+ embedding=cluster_items[0].embedding,
92
+ memory_type=MemoryType.FACTUAL,
93
+ source="consolidator",
94
+ metadata={"consolidated_from": [m.id for m in cluster_items]},
95
+ )
96
+ stored = await self.store.add(new_item)
97
+ summaries.append(stored)
98
+ for m in cluster_items:
99
+ await self.store.update(m.id, status=MemoryStatus.ARCHIVED)
100
+ visited.add(m.id)
101
+ visited.add(stored.id)
102
+ return summaries
103
+
104
+ async def _summarize(self, contents: list[str]) -> str:
105
+ joined = "\n".join(f"- {c}" for c in contents)
106
+ prompt = self.summary_prompt.replace("{memories}", joined)
107
+ response = await self.llm.generate(prompt, temperature=0.0, max_tokens=400)
108
+ return response.strip()
109
+
110
+
111
+ class Pruner:
112
+ """Drop stale memories by strategy.
113
+
114
+ Strategies:
115
+ * ``decay`` — access-count-weighted age; below ``threshold`` is archived.
116
+ * ``age`` — hard cutoff by ``older_than`` timedelta.
117
+ * ``redundancy`` — archive memories whose semantic neighbor count exceeds
118
+ ``max_neighbors`` (i.e., already well-represented).
119
+ """
120
+
121
+ def __init__(self, store: SQLiteStore) -> None:
122
+ self.store = store
123
+
124
+ async def prune(self, strategy: str = "decay", **kwargs: Any) -> int:
125
+ strategy = strategy.lower()
126
+ if strategy == "decay":
127
+ threshold = float(kwargs.get("threshold", 0.1))
128
+ return await self._prune_decay(threshold)
129
+ if strategy == "age":
130
+ older_than = kwargs.get("older_than", timedelta(days=365))
131
+ assert isinstance(older_than, timedelta)
132
+ return await self._prune_age(older_than)
133
+ if strategy == "redundancy":
134
+ semantic = kwargs.get("semantic_graph")
135
+ max_neighbors = int(kwargs.get("max_neighbors", 10))
136
+ if semantic is None:
137
+ return 0
138
+ return await self._prune_redundancy(semantic, max_neighbors)
139
+ raise ValueError(f"Unknown pruning strategy: {strategy}")
140
+
141
+ async def _prune_decay(self, threshold: float) -> int:
142
+ now = datetime.now(tz=timezone.utc)
143
+ pruned = 0
144
+ async for memory in self.store.iter_memories(batch_size=500):
145
+ age_days = max(1.0, (now - memory.created_at).total_seconds() / 86400.0)
146
+ score = (memory.access_count + 1) / age_days
147
+ if score < threshold:
148
+ await self.store.update(memory.id, status=MemoryStatus.ARCHIVED)
149
+ pruned += 1
150
+ return pruned
151
+
152
+ async def _prune_age(self, older_than: timedelta) -> int:
153
+ now = datetime.now(tz=timezone.utc)
154
+ pruned = 0
155
+ async for memory in self.store.iter_memories(batch_size=500):
156
+ if now - memory.created_at > older_than:
157
+ await self.store.update(memory.id, status=MemoryStatus.ARCHIVED)
158
+ pruned += 1
159
+ return pruned
160
+
161
+ async def _prune_redundancy(self, semantic: SemanticGraph, max_neighbors: int) -> int:
162
+ pruned = 0
163
+ async for memory in self.store.iter_memories(batch_size=500):
164
+ neighbors = await semantic.get_neighbors(memory.id, max_results=max_neighbors + 1)
165
+ if len(neighbors) > max_neighbors:
166
+ await self.store.update(memory.id, status=MemoryStatus.ARCHIVED)
167
+ pruned += 1
168
+ return pruned
@@ -0,0 +1,193 @@
1
+ """Memory formation pipeline.
2
+
3
+ Turns a raw conversation or document into a list of ready-to-store
4
+ :class:`~contextdb.core.models.MemoryItem` objects. The steps are:
5
+
6
+ 1. :class:`Segmenter` — split into coherent conversational turns.
7
+ 2. :class:`MemoryExtractor` — LLM pulls atomic facts + entities from a turn.
8
+ 3. :class:`MemoryCompressor` — LLM compresses a cluster into a single summary.
9
+ 4. PII detection + embedding generation happen on the output items.
10
+
11
+ Every step is optional at the call site — the pipeline short-circuits to
12
+ the raw text if the LLM returns nothing usable.
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ import json
18
+ import re
19
+ from typing import TYPE_CHECKING, Any
20
+
21
+ from contextdb.core.models import MemoryItem, MemoryType
22
+
23
+ if TYPE_CHECKING:
24
+ from contextdb.privacy.pii_detector import PIIDetector
25
+ from contextdb.utils.embeddings import EmbeddingProvider
26
+ from contextdb.utils.llm import LLMProvider
27
+
28
+
29
+ _EXTRACT_PROMPT = """Extract atomic facts and named entities from the text.
30
+ Return strict JSON.
31
+
32
+ Schema:
33
+ {"facts": [{"content": "string", "type": "FACTUAL|EXPERIENTIAL", "entities": ["string"]}]}
34
+
35
+ Rules:
36
+ - Each fact must be self-contained (understandable without context).
37
+ - Skip small talk; keep substantive information only.
38
+ - Aim for 1-5 facts per turn.
39
+
40
+ Text: "{text}"
41
+ """
42
+
43
+ _COMPRESS_PROMPT = """Summarize the following related memories into one concise
44
+ statement. Preserve all named entities and dates. Return plain text, no JSON.
45
+
46
+ Memories:
47
+ {memories}
48
+ """
49
+
50
+
51
+ def _safe_json(text: str) -> dict[str, Any]:
52
+ text = text.strip()
53
+ if text.startswith("```"):
54
+ lines = text.splitlines()
55
+ text = "\n".join(line for line in lines if not line.startswith("```"))
56
+ try:
57
+ loaded = json.loads(text)
58
+ return loaded if isinstance(loaded, dict) else {}
59
+ except json.JSONDecodeError:
60
+ start = text.find("{")
61
+ end = text.rfind("}")
62
+ if start != -1 and end != -1 and end > start:
63
+ try:
64
+ loaded = json.loads(text[start : end + 1])
65
+ return loaded if isinstance(loaded, dict) else {}
66
+ except json.JSONDecodeError:
67
+ return {}
68
+ return {}
69
+
70
+
71
+ class Segmenter:
72
+ """Split raw text into turns / coherent chunks.
73
+
74
+ The baseline rule: newlines separate turns, and speaker prefixes
75
+ (``User:``, ``Agent:``) are preserved. Anything shorter than ``min_chars``
76
+ is merged into the next chunk so we don't emit fragments.
77
+ """
78
+
79
+ def __init__(self, min_chars: int = 20) -> None:
80
+ self.min_chars = min_chars
81
+
82
+ def segment(self, text: str) -> list[str]:
83
+ raw = [
84
+ chunk.strip()
85
+ for chunk in re.split(r"\n{2,}|(?<=[.!?])\s{2,}", text)
86
+ if chunk.strip()
87
+ ]
88
+ merged: list[str] = []
89
+ buffer = ""
90
+ for chunk in raw:
91
+ candidate = f"{buffer} {chunk}".strip() if buffer else chunk
92
+ if len(candidate) < self.min_chars:
93
+ buffer = candidate
94
+ continue
95
+ merged.append(candidate)
96
+ buffer = ""
97
+ if buffer:
98
+ if merged:
99
+ merged[-1] = f"{merged[-1]} {buffer}".strip()
100
+ else:
101
+ merged.append(buffer)
102
+ return merged
103
+
104
+
105
+ class MemoryExtractor:
106
+ """LLM-driven fact + entity extraction per turn."""
107
+
108
+ def __init__(self, llm: LLMProvider) -> None:
109
+ self.llm = llm
110
+
111
+ async def extract(self, turn: str) -> list[dict[str, Any]]:
112
+ response = await self.llm.generate(_EXTRACT_PROMPT.replace("{text}", turn))
113
+ payload = _safe_json(response)
114
+ out: list[dict[str, Any]] = []
115
+ for raw in payload.get("facts", []) or []:
116
+ content = str(raw.get("content", "")).strip()
117
+ if not content:
118
+ continue
119
+ mem_type = str(raw.get("type", "FACTUAL")).upper()
120
+ if mem_type not in {"FACTUAL", "EXPERIENTIAL", "WORKING"}:
121
+ mem_type = "FACTUAL"
122
+ entities = [str(e).strip() for e in raw.get("entities", []) or [] if e]
123
+ out.append({"content": content, "memory_type": mem_type, "entities": entities})
124
+ return out
125
+
126
+
127
+ class MemoryCompressor:
128
+ """LLM-driven cluster summarization.
129
+
130
+ Given a list of memory contents, produce a single condensed statement
131
+ that preserves entities and temporal markers. If the LLM returns an empty
132
+ string, we fall back to naïve concatenation so the caller never loses
133
+ data.
134
+ """
135
+
136
+ def __init__(self, llm: LLMProvider) -> None:
137
+ self.llm = llm
138
+
139
+ async def compress(self, memories: list[str]) -> str:
140
+ if not memories:
141
+ return ""
142
+ if len(memories) == 1:
143
+ return memories[0]
144
+ joined = "\n".join(f"- {m}" for m in memories)
145
+ response = await self.llm.generate(_COMPRESS_PROMPT.replace("{memories}", joined))
146
+ summary = response.strip()
147
+ return summary or " | ".join(memories)
148
+
149
+
150
+ class FormationPipeline:
151
+ """Glue the formation steps into a single async entry point."""
152
+
153
+ def __init__(
154
+ self,
155
+ segmenter: Segmenter,
156
+ extractor: MemoryExtractor,
157
+ compressor: MemoryCompressor,
158
+ pii: PIIDetector,
159
+ embedder: EmbeddingProvider,
160
+ ) -> None:
161
+ self.segmenter = segmenter
162
+ self.extractor = extractor
163
+ self.compressor = compressor
164
+ self.pii = pii
165
+ self.embedder = embedder
166
+
167
+ async def process(self, text: str, source: str = "") -> list[MemoryItem]:
168
+ turns = self.segmenter.segment(text)
169
+ all_facts: list[dict[str, Any]] = []
170
+ for turn in turns:
171
+ facts = await self.extractor.extract(turn)
172
+ if not facts:
173
+ # Fallback: store the turn verbatim as a FACTUAL memory.
174
+ facts = [{"content": turn, "memory_type": "FACTUAL", "entities": []}]
175
+ all_facts.extend(facts)
176
+
177
+ items: list[MemoryItem] = []
178
+ contents = [fact["content"] for fact in all_facts]
179
+ embeddings = await self.embedder.embed(contents) if contents else []
180
+ for fact, embedding in zip(all_facts, embeddings, strict=False):
181
+ content = fact["content"]
182
+ processed, pii_annotations = self.pii.process(content)
183
+ items.append(
184
+ MemoryItem(
185
+ content=processed,
186
+ embedding=embedding,
187
+ memory_type=MemoryType(fact["memory_type"]),
188
+ source=source,
189
+ pii_annotations=pii_annotations,
190
+ entity_mentions=list(fact.get("entities", [])),
191
+ )
192
+ )
193
+ return items
@@ -0,0 +1,130 @@
1
+ """Multi-graph retrieval — query classification + reciprocal-rank fusion.
2
+
3
+ The engine runs the query against each configured graph (plus a raw vector
4
+ search against the store) and fuses results using Reciprocal Rank Fusion
5
+ (Cormack et al., 2009) with the classic ``k=60`` smoothing. RRF is
6
+ parameter-light, ignores raw score scales, and tends to beat linear
7
+ combination for heterogeneous retrievers — ideal when one retriever scores in
8
+ cosine space and another in edge-weight space.
9
+
10
+ The :class:`QueryClassifier` is intentionally rule-based. An LLM classifier
11
+ would be more accurate but would add latency to every query; the regex
12
+ heuristics below get ~80% of the signal for zero cost.
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ import re
18
+ from typing import TYPE_CHECKING
19
+
20
+ from contextdb.core.models import MemoryItem
21
+
22
+ if TYPE_CHECKING:
23
+ from contextdb.graphs.base import BaseGraph
24
+ from contextdb.store.sqlite_store import SQLiteStore
25
+
26
+
27
+ _TEMPORAL_MARKERS = re.compile(
28
+ r"\b(when|before|after|yesterday|today|tomorrow|last|next|since|until|during|ago)\b",
29
+ re.IGNORECASE,
30
+ )
31
+ _CAUSAL_MARKERS = re.compile(
32
+ r"\b(why|because|caused|leads? to|due to|resulted? in|reason|so that)\b",
33
+ re.IGNORECASE,
34
+ )
35
+ _ENTITY_MARKERS = re.compile(
36
+ r"\b(who|whose|which person|what company|what product)\b",
37
+ re.IGNORECASE,
38
+ )
39
+
40
+
41
+ class QueryClassifier:
42
+ """Classify a natural-language query into graph-weighting hints.
43
+
44
+ Returns a ``dict[graph_name, weight]`` that sums roughly to 1.0. Callers
45
+ use these as mixing weights when fusing per-graph rankings.
46
+ """
47
+
48
+ def classify(self, query: str) -> dict[str, float]:
49
+ weights: dict[str, float] = {"semantic": 1.0}
50
+ if _TEMPORAL_MARKERS.search(query):
51
+ weights["temporal"] = 1.2
52
+ if _CAUSAL_MARKERS.search(query):
53
+ weights["causal"] = 1.4
54
+ if _ENTITY_MARKERS.search(query):
55
+ weights["entity"] = 1.1
56
+ total = sum(weights.values())
57
+ return {k: v / total for k, v in weights.items()}
58
+
59
+
60
+ class RetrievalFuser:
61
+ """Reciprocal Rank Fusion over per-graph candidate lists."""
62
+
63
+ def __init__(self, k: int = 60) -> None:
64
+ self.k = k
65
+
66
+ def fuse(
67
+ self,
68
+ rankings: dict[str, list[tuple[str, float]]],
69
+ weights: dict[str, float],
70
+ ) -> list[tuple[str, float]]:
71
+ scores: dict[str, float] = {}
72
+ for graph_name, ranking in rankings.items():
73
+ w = weights.get(graph_name, 0.0)
74
+ if w == 0.0 or not ranking:
75
+ continue
76
+ for rank, (memory_id, _) in enumerate(ranking, start=1):
77
+ scores[memory_id] = scores.get(memory_id, 0.0) + w * (1.0 / (self.k + rank))
78
+ return sorted(scores.items(), key=lambda kv: kv[1], reverse=True)
79
+
80
+
81
+ class RetrievalEngine:
82
+ """Coordinate vector + graph retrieval and return ranked memories."""
83
+
84
+ def __init__(
85
+ self,
86
+ store: SQLiteStore,
87
+ graphs: dict[str, BaseGraph],
88
+ classifier: QueryClassifier,
89
+ fuser: RetrievalFuser,
90
+ ) -> None:
91
+ self.store = store
92
+ self.graphs = graphs
93
+ self.classifier = classifier
94
+ self.fuser = fuser
95
+
96
+ async def search(
97
+ self,
98
+ query: str,
99
+ query_embedding: list[float],
100
+ top_k: int = 10,
101
+ ) -> list[MemoryItem]:
102
+ weights = self.classifier.classify(query)
103
+ seed_items = await self.store.search_by_embedding(query_embedding, top_k=top_k * 2)
104
+ semantic_ranking = [(item.id, 1.0 / (rank + 1)) for rank, item in enumerate(seed_items)]
105
+ rankings: dict[str, list[tuple[str, float]]] = {"semantic": semantic_ranking}
106
+
107
+ seed_ids = [item.id for item in seed_items[: max(1, top_k)]]
108
+ for name, graph in self.graphs.items():
109
+ if name == "semantic":
110
+ continue
111
+ if weights.get(name, 0.0) <= 0.0:
112
+ continue
113
+ expanded: dict[str, float] = {}
114
+ for sid in seed_ids:
115
+ neighbors = await graph.get_neighbors(sid, max_results=top_k)
116
+ for nid, weight in neighbors:
117
+ expanded[nid] = max(expanded.get(nid, 0.0), weight)
118
+ rankings[name] = sorted(expanded.items(), key=lambda kv: kv[1], reverse=True)
119
+
120
+ fused = self.fuser.fuse(rankings, weights)
121
+ ordered_ids = [mid for mid, _ in fused[:top_k]]
122
+ if not ordered_ids:
123
+ return seed_items[:top_k]
124
+
125
+ items: list[MemoryItem] = []
126
+ for mid in ordered_ids:
127
+ item = await self.store.get_raw(mid)
128
+ if item is not None:
129
+ items.append(item)
130
+ return items
@@ -0,0 +1,17 @@
1
+ """Graph views over memories: semantic, temporal, causal, entity."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from contextdb.graphs.base import BaseGraph
6
+ from contextdb.graphs.causal import CausalGraph
7
+ from contextdb.graphs.entity import EntityGraph
8
+ from contextdb.graphs.semantic import SemanticGraph
9
+ from contextdb.graphs.temporal import TemporalGraph
10
+
11
+ __all__ = [
12
+ "BaseGraph",
13
+ "CausalGraph",
14
+ "EntityGraph",
15
+ "SemanticGraph",
16
+ "TemporalGraph",
17
+ ]