mcp-agensgraph-memory 0.2.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.
- mcp_agensgraph_memory/__init__.py +55 -0
- mcp_agensgraph_memory/agensgraph_memory.py +546 -0
- mcp_agensgraph_memory/server.py +584 -0
- mcp_agensgraph_memory/utils.py +27 -0
- mcp_agensgraph_memory-0.2.0.dist-info/METADATA +301 -0
- mcp_agensgraph_memory-0.2.0.dist-info/RECORD +10 -0
- mcp_agensgraph_memory-0.2.0.dist-info/WHEEL +4 -0
- mcp_agensgraph_memory-0.2.0.dist-info/entry_points.txt +2 -0
- mcp_agensgraph_memory-0.2.0.dist-info/licenses/LICENSE +201 -0
- mcp_agensgraph_memory-0.2.0.dist-info/licenses/NOTICE +35 -0
|
@@ -0,0 +1,55 @@
|
|
|
1
|
+
import argparse
|
|
2
|
+
import asyncio
|
|
3
|
+
import logging
|
|
4
|
+
|
|
5
|
+
from . import server
|
|
6
|
+
from .utils import process_config
|
|
7
|
+
|
|
8
|
+
logger = logging.getLogger("mcp_agensgraph_memory")
|
|
9
|
+
logger.setLevel(logging.INFO)
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def main():
|
|
13
|
+
"""Main entry point for the package."""
|
|
14
|
+
parser = argparse.ArgumentParser(description="AgensGraph Memory MCP Server")
|
|
15
|
+
parser.add_argument(
|
|
16
|
+
"--db-url",
|
|
17
|
+
default=None,
|
|
18
|
+
help="AgensGraph connection URL (postgresql://host:port)",
|
|
19
|
+
)
|
|
20
|
+
parser.add_argument("--username", default=None, help="AgensGraph username")
|
|
21
|
+
parser.add_argument("--password", default=None, help="AgensGraph password")
|
|
22
|
+
parser.add_argument("--database", default=None, help="AgensGraph database name")
|
|
23
|
+
parser.add_argument("--graphname", default=None, help="AgensGraph graph name")
|
|
24
|
+
parser.add_argument("--namespace", default=None, help="Tool namespace prefix")
|
|
25
|
+
parser.add_argument(
|
|
26
|
+
"--transport", default=None, help="Transport type (stdio, sse, http)"
|
|
27
|
+
)
|
|
28
|
+
parser.add_argument(
|
|
29
|
+
"--server-host", default=None, help="HTTP host (default: 127.0.0.1)"
|
|
30
|
+
)
|
|
31
|
+
parser.add_argument(
|
|
32
|
+
"--server-port", type=int, default=None, help="HTTP port (default: 8000)"
|
|
33
|
+
)
|
|
34
|
+
parser.add_argument(
|
|
35
|
+
"--server-path", default=None, help="HTTP path (default: /mcp/)"
|
|
36
|
+
)
|
|
37
|
+
parser.add_argument(
|
|
38
|
+
"--allow-origins",
|
|
39
|
+
default=None,
|
|
40
|
+
help="Comma-separated list of allowed CORS origins",
|
|
41
|
+
)
|
|
42
|
+
parser.add_argument(
|
|
43
|
+
"--allowed-hosts",
|
|
44
|
+
default=None,
|
|
45
|
+
help="Comma-separated list of allowed hosts for DNS rebinding protection",
|
|
46
|
+
)
|
|
47
|
+
|
|
48
|
+
args = parser.parse_args()
|
|
49
|
+
|
|
50
|
+
config = process_config(args)
|
|
51
|
+
asyncio.run(server.main(**config))
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
# Optionally expose other important items at package level
|
|
55
|
+
__all__ = ["main", "server"]
|
|
@@ -0,0 +1,546 @@
|
|
|
1
|
+
import logging
|
|
2
|
+
from typing import Any, Dict, List
|
|
3
|
+
|
|
4
|
+
from psycopg import AsyncConnection, sql
|
|
5
|
+
from psycopg.rows import namedtuple_row
|
|
6
|
+
from psycopg.types.json import Jsonb
|
|
7
|
+
from psycopg_pool import AsyncConnectionPool
|
|
8
|
+
from pydantic import BaseModel, Field
|
|
9
|
+
|
|
10
|
+
from mcp_agensgraph_common.connection import get_pool_connection
|
|
11
|
+
from mcp_agensgraph_common.safety import quote_label
|
|
12
|
+
|
|
13
|
+
# Set up logging
|
|
14
|
+
logger = logging.getLogger("mcp_agensgraph_memory")
|
|
15
|
+
logger.setLevel(logging.INFO)
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
# Models for our knowledge graph
|
|
19
|
+
class Entity(BaseModel):
|
|
20
|
+
"""Represents a memory entity in the knowledge graph.
|
|
21
|
+
|
|
22
|
+
Example:
|
|
23
|
+
{
|
|
24
|
+
"name": "John Smith",
|
|
25
|
+
"type": "person",
|
|
26
|
+
"observations": ["Works at SKAI Worldwide", "Lives in San Francisco", "Expert in graph databases"]
|
|
27
|
+
}
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
name: str = Field(
|
|
31
|
+
description="Unique identifier/name for the entity. Should be descriptive and specific.",
|
|
32
|
+
min_length=1,
|
|
33
|
+
examples=["John Smith", "SKAI Worldwide Inc", "San Francisco"],
|
|
34
|
+
)
|
|
35
|
+
type: str = Field(
|
|
36
|
+
description="Category or classification of the entity. Common types: 'person', 'company', 'location', 'concept', 'event'",
|
|
37
|
+
min_length=1,
|
|
38
|
+
examples=["person", "company", "location", "concept", "event"],
|
|
39
|
+
)
|
|
40
|
+
observations: List[str] = Field(
|
|
41
|
+
description="List of facts, observations, or notes about this entity. Each observation should be a complete, standalone fact.",
|
|
42
|
+
examples=[
|
|
43
|
+
["Works at SKAI Worldwide", "Lives in San Francisco"],
|
|
44
|
+
["Headquartered in Sweden", "Graph database company"],
|
|
45
|
+
],
|
|
46
|
+
)
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
class Relation(BaseModel):
|
|
50
|
+
"""Represents a relationship between two entities in the knowledge graph.
|
|
51
|
+
|
|
52
|
+
Example:
|
|
53
|
+
{
|
|
54
|
+
"source": "John Smith",
|
|
55
|
+
"target": "SKAI Worldwide Inc",
|
|
56
|
+
"relationType": "WORKS_AT"
|
|
57
|
+
}
|
|
58
|
+
"""
|
|
59
|
+
|
|
60
|
+
source: str = Field(
|
|
61
|
+
description="Name of the source entity (must match an existing entity name exactly)",
|
|
62
|
+
min_length=1,
|
|
63
|
+
examples=["John Smith", "SKAI Worldwide Inc"],
|
|
64
|
+
)
|
|
65
|
+
target: str = Field(
|
|
66
|
+
description="Name of the target entity (must match an existing entity name exactly)",
|
|
67
|
+
min_length=1,
|
|
68
|
+
examples=["SKAI Worldwide Inc", "San Francisco"],
|
|
69
|
+
)
|
|
70
|
+
relationType: str = Field(
|
|
71
|
+
description="Type of relationship between source and target. Use descriptive, uppercase names with underscores.",
|
|
72
|
+
min_length=1,
|
|
73
|
+
examples=["WORKS_AT", "LIVES_IN", "MANAGES", "COLLABORATES_WITH", "LOCATED_IN"],
|
|
74
|
+
)
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
class KnowledgeGraph(BaseModel):
|
|
78
|
+
"""Complete knowledge graph containing entities and their relationships."""
|
|
79
|
+
|
|
80
|
+
entities: List[Entity] = Field(
|
|
81
|
+
description="List of all entities in the knowledge graph", default=[]
|
|
82
|
+
)
|
|
83
|
+
relations: List[Relation] = Field(
|
|
84
|
+
description="List of all relationships between entities", default=[]
|
|
85
|
+
)
|
|
86
|
+
truncated: bool = Field(
|
|
87
|
+
default=False,
|
|
88
|
+
description=(
|
|
89
|
+
"True if the entity list was capped by the limit. Narrow with "
|
|
90
|
+
"search_memories or request a higher limit to see more."
|
|
91
|
+
),
|
|
92
|
+
)
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
class ObservationAddition(BaseModel):
|
|
96
|
+
"""Request to add new observations to an existing entity.
|
|
97
|
+
|
|
98
|
+
Example:
|
|
99
|
+
{
|
|
100
|
+
"entityName": "John Smith",
|
|
101
|
+
"observations": ["Recently promoted to Senior Engineer", "Speaks fluent German"]
|
|
102
|
+
}
|
|
103
|
+
"""
|
|
104
|
+
|
|
105
|
+
entityName: str = Field(
|
|
106
|
+
description="Exact name of the existing entity to add observations to",
|
|
107
|
+
min_length=1,
|
|
108
|
+
examples=["John Smith", "SKAI Worldwide Inc"],
|
|
109
|
+
)
|
|
110
|
+
observations: List[str] = Field(
|
|
111
|
+
description="New observations/facts to add to the entity. Each should be unique and informative.",
|
|
112
|
+
min_length=1,
|
|
113
|
+
)
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
class ObservationDeletion(BaseModel):
|
|
117
|
+
"""Request to delete specific observations from an existing entity.
|
|
118
|
+
|
|
119
|
+
Example:
|
|
120
|
+
{
|
|
121
|
+
"entityName": "John Smith",
|
|
122
|
+
"observations": ["Old job title", "Outdated contact info"]
|
|
123
|
+
}
|
|
124
|
+
"""
|
|
125
|
+
|
|
126
|
+
entityName: str = Field(
|
|
127
|
+
description="Exact name of the existing entity to remove observations from",
|
|
128
|
+
min_length=1,
|
|
129
|
+
examples=["John Smith", "SKAI Worldwide Inc"],
|
|
130
|
+
)
|
|
131
|
+
observations: List[str] = Field(
|
|
132
|
+
description="Exact observation texts to delete from the entity (must match existing observations exactly)",
|
|
133
|
+
min_length=1,
|
|
134
|
+
)
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
class AgensGraphMemory:
|
|
138
|
+
def __init__(self, connection_pool: AsyncConnectionPool, graphname: str):
|
|
139
|
+
self.pool = connection_pool
|
|
140
|
+
self.graphname = graphname
|
|
141
|
+
|
|
142
|
+
async def _execute_cypher(
|
|
143
|
+
self, conn: AsyncConnection, cypher_query: str, params: dict = None
|
|
144
|
+
):
|
|
145
|
+
"""Execute a Cypher query within AgensGraph."""
|
|
146
|
+
async with conn.cursor(row_factory=namedtuple_row) as cursor:
|
|
147
|
+
# Set graph path
|
|
148
|
+
await cursor.execute(
|
|
149
|
+
sql.SQL("SET graph_path = {}").format(sql.Identifier(self.graphname))
|
|
150
|
+
)
|
|
151
|
+
|
|
152
|
+
# Execute the Cypher query
|
|
153
|
+
if params:
|
|
154
|
+
await cursor.execute(cypher_query, params)
|
|
155
|
+
else:
|
|
156
|
+
await cursor.execute(cypher_query)
|
|
157
|
+
|
|
158
|
+
# Try to fetch results
|
|
159
|
+
try:
|
|
160
|
+
results = await cursor.fetchall()
|
|
161
|
+
return results
|
|
162
|
+
except Exception:
|
|
163
|
+
# Query might not return results (INSERT, DELETE, etc.)
|
|
164
|
+
return []
|
|
165
|
+
|
|
166
|
+
async def create_fulltext_index(self):
|
|
167
|
+
"""Create a fulltext search index for entities if it doesn't exist.
|
|
168
|
+
|
|
169
|
+
Uses PostgreSQL's text search (tsvector) for fulltext search capabilities.
|
|
170
|
+
Creates a GIN property index on the tsvector for efficient searching.
|
|
171
|
+
Also ensures the Memory VLABEL exists.
|
|
172
|
+
"""
|
|
173
|
+
try:
|
|
174
|
+
async with get_pool_connection(self.pool) as conn:
|
|
175
|
+
async with conn.cursor(row_factory=namedtuple_row) as cursor:
|
|
176
|
+
# Set graph path
|
|
177
|
+
await cursor.execute(
|
|
178
|
+
sql.SQL("SET graph_path = {}").format(sql.Identifier(self.graphname))
|
|
179
|
+
)
|
|
180
|
+
|
|
181
|
+
# Ensure Memory VLABEL exists
|
|
182
|
+
await cursor.execute('CREATE VLABEL IF NOT EXISTS "Memory"')
|
|
183
|
+
|
|
184
|
+
# Create GIN property index on tsvector for fulltext search
|
|
185
|
+
# This indexes name, type, and all observations combined
|
|
186
|
+
await cursor.execute("""
|
|
187
|
+
CREATE PROPERTY INDEX IF NOT EXISTS memory_fulltext_idx
|
|
188
|
+
ON "Memory"
|
|
189
|
+
USING gin
|
|
190
|
+
(
|
|
191
|
+
(
|
|
192
|
+
setweight(to_tsvector('english', coalesce(name, '')), 'A') ||
|
|
193
|
+
setweight(to_tsvector('english', coalesce(type, '')), 'B') ||
|
|
194
|
+
setweight(to_tsvector('english', coalesce(jsonb_to_string(observations, ' '), '')), 'C')
|
|
195
|
+
)
|
|
196
|
+
)
|
|
197
|
+
""")
|
|
198
|
+
await conn.commit()
|
|
199
|
+
logger.info(
|
|
200
|
+
"Created Memory VLABEL and fulltext search property index using tsvector"
|
|
201
|
+
)
|
|
202
|
+
except Exception as e:
|
|
203
|
+
# Index might already exist, which is fine
|
|
204
|
+
logger.debug(f"Fulltext property index creation: {e}")
|
|
205
|
+
|
|
206
|
+
async def load_graph(self, filter_query: str = None, limit: int = None):
|
|
207
|
+
"""Load the knowledge graph from AgensGraph.
|
|
208
|
+
|
|
209
|
+
If ``limit`` is set, at most that many entities are returned (capped at the
|
|
210
|
+
database with ``LIMIT`` so a large memory can't flood the caller's context);
|
|
211
|
+
the returned graph's ``truncated`` flag indicates whether more exist.
|
|
212
|
+
Relations reference entities by name, so the capped graph stays coherent.
|
|
213
|
+
"""
|
|
214
|
+
logger.info("Loading knowledge graph from AgensGraph")
|
|
215
|
+
|
|
216
|
+
async with get_pool_connection(self.pool) as conn:
|
|
217
|
+
# Build the filter condition using PostgreSQL fulltext search
|
|
218
|
+
if filter_query and filter_query != "*":
|
|
219
|
+
# Use tsvector and tsquery for fulltext search
|
|
220
|
+
# Searches across name, type, and observations with weights
|
|
221
|
+
filter_condition = """
|
|
222
|
+
WHERE (
|
|
223
|
+
setweight(to_tsvector('english', coalesce(entity.name, '')), 'A') ||
|
|
224
|
+
setweight(to_tsvector('english', coalesce(entity.type, '')), 'B') ||
|
|
225
|
+
setweight(to_tsvector('english', coalesce(jsonb_to_string(entity.observations, ' '), '')), 'C')
|
|
226
|
+
) @@ plainto_tsquery('english', %(query)s)
|
|
227
|
+
"""
|
|
228
|
+
params = {"query": filter_query}
|
|
229
|
+
else:
|
|
230
|
+
filter_condition = ""
|
|
231
|
+
params = {}
|
|
232
|
+
|
|
233
|
+
# Fetch one extra row to detect truncation when a limit is applied.
|
|
234
|
+
# limit is a validated int, so inlining it is injection-safe.
|
|
235
|
+
limit_clause = ""
|
|
236
|
+
if limit is not None:
|
|
237
|
+
limit = max(1, int(limit))
|
|
238
|
+
limit_clause = f"ORDER BY entity.name LIMIT {limit + 1}"
|
|
239
|
+
|
|
240
|
+
entity_query = f"""
|
|
241
|
+
MATCH (entity:"Memory")
|
|
242
|
+
{filter_condition}
|
|
243
|
+
RETURN entity.name AS name, entity.type AS type, entity.observations AS observations
|
|
244
|
+
{limit_clause}
|
|
245
|
+
"""
|
|
246
|
+
entities_data = await self._execute_cypher(conn, entity_query, params)
|
|
247
|
+
|
|
248
|
+
truncated = limit is not None and len(entities_data) > limit
|
|
249
|
+
if truncated:
|
|
250
|
+
entities_data = entities_data[:limit]
|
|
251
|
+
|
|
252
|
+
entities = []
|
|
253
|
+
entity_names = []
|
|
254
|
+
for record in entities_data:
|
|
255
|
+
entities.append(
|
|
256
|
+
Entity(
|
|
257
|
+
name=record.name,
|
|
258
|
+
type=record.type,
|
|
259
|
+
observations=record.observations or [],
|
|
260
|
+
)
|
|
261
|
+
)
|
|
262
|
+
entity_names.append(record.name)
|
|
263
|
+
|
|
264
|
+
# Query to get all relationships for these entities
|
|
265
|
+
relations = []
|
|
266
|
+
if entity_names:
|
|
267
|
+
rel_query = """
|
|
268
|
+
MATCH (source:"Memory")-[r]->(target:"Memory")
|
|
269
|
+
WHERE source.name IN %(names)s OR target.name IN %(names)s
|
|
270
|
+
RETURN source.name AS source, target.name AS target, label(r) AS "relationType"
|
|
271
|
+
"""
|
|
272
|
+
|
|
273
|
+
relations_data = await self._execute_cypher(
|
|
274
|
+
conn, rel_query, {"names": Jsonb(entity_names)}
|
|
275
|
+
)
|
|
276
|
+
|
|
277
|
+
for record in relations_data:
|
|
278
|
+
relations.append(
|
|
279
|
+
Relation(
|
|
280
|
+
source=record.source,
|
|
281
|
+
target=record.target,
|
|
282
|
+
relationType=record.relationType,
|
|
283
|
+
)
|
|
284
|
+
)
|
|
285
|
+
|
|
286
|
+
await conn.commit()
|
|
287
|
+
|
|
288
|
+
logger.debug(f"Loaded entities: {entities}")
|
|
289
|
+
logger.debug(f"Loaded relations: {relations}")
|
|
290
|
+
|
|
291
|
+
return KnowledgeGraph(
|
|
292
|
+
entities=entities, relations=relations, truncated=truncated
|
|
293
|
+
)
|
|
294
|
+
|
|
295
|
+
async def create_entities(self, entities: List[Entity]) -> List[Entity]:
|
|
296
|
+
"""Create multiple new entities in the knowledge graph."""
|
|
297
|
+
logger.info(f"Creating {len(entities)} entities")
|
|
298
|
+
|
|
299
|
+
async with get_pool_connection(self.pool) as conn:
|
|
300
|
+
for entity in entities:
|
|
301
|
+
# Create/update the entity
|
|
302
|
+
# Note: We store the type as a property, not as a separate label
|
|
303
|
+
query = """
|
|
304
|
+
MERGE (e:"Memory" {name: %(name)s})
|
|
305
|
+
SET e.type = %(type)s, e.observations = %(observations)s
|
|
306
|
+
"""
|
|
307
|
+
|
|
308
|
+
await self._execute_cypher(
|
|
309
|
+
conn,
|
|
310
|
+
query,
|
|
311
|
+
{
|
|
312
|
+
"name": Jsonb(entity.name),
|
|
313
|
+
"type": Jsonb(entity.type),
|
|
314
|
+
"observations": Jsonb(entity.observations),
|
|
315
|
+
},
|
|
316
|
+
)
|
|
317
|
+
|
|
318
|
+
await conn.commit()
|
|
319
|
+
|
|
320
|
+
return entities
|
|
321
|
+
|
|
322
|
+
async def create_relations(self, relations: List[Relation]) -> List[Relation]:
|
|
323
|
+
"""Create multiple new relations between entities."""
|
|
324
|
+
logger.info(f"Creating {len(relations)} relations")
|
|
325
|
+
|
|
326
|
+
async with get_pool_connection(self.pool) as conn:
|
|
327
|
+
for relation in relations:
|
|
328
|
+
# create the relationship
|
|
329
|
+
# Relationship types cannot be parameterized in Cypher, so the
|
|
330
|
+
# (client-supplied) type is validated + identifier-quoted rather
|
|
331
|
+
# than interpolated raw.
|
|
332
|
+
rel_type = quote_label(relation.relationType)
|
|
333
|
+
query = f"""
|
|
334
|
+
MATCH (fromNode:"Memory"), (toNode:"Memory")
|
|
335
|
+
WHERE fromNode.name = %(source)s AND toNode.name = %(target)s
|
|
336
|
+
MERGE (fromNode)-[r:{rel_type}]->(toNode)
|
|
337
|
+
"""
|
|
338
|
+
|
|
339
|
+
await self._execute_cypher(
|
|
340
|
+
conn,
|
|
341
|
+
query,
|
|
342
|
+
{
|
|
343
|
+
"source": Jsonb(relation.source),
|
|
344
|
+
"target": Jsonb(relation.target),
|
|
345
|
+
},
|
|
346
|
+
)
|
|
347
|
+
|
|
348
|
+
await conn.commit()
|
|
349
|
+
|
|
350
|
+
return relations
|
|
351
|
+
|
|
352
|
+
async def add_observations(
|
|
353
|
+
self, observations: List[ObservationAddition]
|
|
354
|
+
) -> List[Dict[str, Any]]:
|
|
355
|
+
"""Add new observations to existing entities."""
|
|
356
|
+
logger.info(f"Adding observations to {len(observations)} entities")
|
|
357
|
+
|
|
358
|
+
results = []
|
|
359
|
+
async with get_pool_connection(self.pool) as conn:
|
|
360
|
+
for obs in observations:
|
|
361
|
+
# Get existing observations
|
|
362
|
+
get_query = """
|
|
363
|
+
MATCH (e:"Memory" {name: %(name)s})
|
|
364
|
+
RETURN e.observations AS observations
|
|
365
|
+
"""
|
|
366
|
+
|
|
367
|
+
existing_data = await self._execute_cypher(
|
|
368
|
+
conn, get_query, {"name": Jsonb(obs.entityName)}
|
|
369
|
+
)
|
|
370
|
+
|
|
371
|
+
if existing_data:
|
|
372
|
+
existing_obs = existing_data[0].observations or []
|
|
373
|
+
# Filter out observations that already exist
|
|
374
|
+
new_obs = [o for o in obs.observations if o not in existing_obs]
|
|
375
|
+
|
|
376
|
+
if new_obs:
|
|
377
|
+
# Update with new observations using Cypher list concatenation
|
|
378
|
+
update_query = """
|
|
379
|
+
MATCH (e:"Memory" {name: %(name)s})
|
|
380
|
+
SET e.observations = coalesce(e.observations, []) + %(new_obs)s
|
|
381
|
+
"""
|
|
382
|
+
|
|
383
|
+
await self._execute_cypher(
|
|
384
|
+
conn,
|
|
385
|
+
update_query,
|
|
386
|
+
{"name": Jsonb(obs.entityName), "new_obs": Jsonb(new_obs)},
|
|
387
|
+
)
|
|
388
|
+
|
|
389
|
+
results.append(
|
|
390
|
+
{"entityName": obs.entityName, "addedObservations": new_obs}
|
|
391
|
+
)
|
|
392
|
+
else:
|
|
393
|
+
results.append(
|
|
394
|
+
{"entityName": obs.entityName, "addedObservations": []}
|
|
395
|
+
)
|
|
396
|
+
|
|
397
|
+
await conn.commit()
|
|
398
|
+
|
|
399
|
+
return results
|
|
400
|
+
|
|
401
|
+
async def delete_entities(self, entity_names: List[str]) -> None:
|
|
402
|
+
"""Delete multiple entities and their associated relations."""
|
|
403
|
+
logger.info(f"Deleting {len(entity_names)} entities")
|
|
404
|
+
|
|
405
|
+
async with get_pool_connection(self.pool) as conn:
|
|
406
|
+
for name in entity_names:
|
|
407
|
+
query = """
|
|
408
|
+
MATCH (e:"Memory" {name: %(name)s})
|
|
409
|
+
DETACH DELETE e
|
|
410
|
+
"""
|
|
411
|
+
|
|
412
|
+
await self._execute_cypher(conn, query, {"name": Jsonb(name)})
|
|
413
|
+
|
|
414
|
+
await conn.commit()
|
|
415
|
+
|
|
416
|
+
logger.info(f"Successfully deleted {len(entity_names)} entities")
|
|
417
|
+
|
|
418
|
+
async def delete_observations(self, deletions: List[ObservationDeletion]) -> None:
|
|
419
|
+
"""Delete specific observations from entities."""
|
|
420
|
+
logger.info(f"Deleting observations from {len(deletions)} entities")
|
|
421
|
+
|
|
422
|
+
async with get_pool_connection(self.pool) as conn:
|
|
423
|
+
for deletion in deletions:
|
|
424
|
+
# Get existing observations
|
|
425
|
+
get_query = """
|
|
426
|
+
MATCH (e:"Memory" {name: %(name)s})
|
|
427
|
+
RETURN e.observations AS observations
|
|
428
|
+
"""
|
|
429
|
+
|
|
430
|
+
existing_data = await self._execute_cypher(
|
|
431
|
+
conn, get_query, {"name": Jsonb(deletion.entityName)}
|
|
432
|
+
)
|
|
433
|
+
|
|
434
|
+
if existing_data:
|
|
435
|
+
existing_obs = existing_data[0].observations or []
|
|
436
|
+
# Filter out observations to delete
|
|
437
|
+
remaining_obs = [
|
|
438
|
+
o for o in existing_obs if o not in deletion.observations
|
|
439
|
+
]
|
|
440
|
+
|
|
441
|
+
# Update with remaining observations (no casting needed in Cypher)
|
|
442
|
+
update_query = """
|
|
443
|
+
MATCH (e:"Memory" {name: %(name)s})
|
|
444
|
+
SET e.observations = %(remaining_obs)s
|
|
445
|
+
"""
|
|
446
|
+
|
|
447
|
+
await self._execute_cypher(
|
|
448
|
+
conn,
|
|
449
|
+
update_query,
|
|
450
|
+
{
|
|
451
|
+
"name": Jsonb(deletion.entityName),
|
|
452
|
+
"remaining_obs": Jsonb(remaining_obs),
|
|
453
|
+
},
|
|
454
|
+
)
|
|
455
|
+
|
|
456
|
+
await conn.commit()
|
|
457
|
+
|
|
458
|
+
logger.info(f"Successfully deleted observations from {len(deletions)} entities")
|
|
459
|
+
|
|
460
|
+
async def delete_relations(self, relations: List[Relation]) -> None:
|
|
461
|
+
"""Delete multiple relations from the graph."""
|
|
462
|
+
logger.info(f"Deleting {len(relations)} relations")
|
|
463
|
+
|
|
464
|
+
async with get_pool_connection(self.pool) as conn:
|
|
465
|
+
for relation in relations:
|
|
466
|
+
rel_type = quote_label(relation.relationType)
|
|
467
|
+
query = f"""
|
|
468
|
+
MATCH (source:"Memory")-[r:{rel_type}]->(target:"Memory")
|
|
469
|
+
WHERE source.name = %(source)s AND target.name = %(target)s
|
|
470
|
+
DELETE r
|
|
471
|
+
"""
|
|
472
|
+
|
|
473
|
+
await self._execute_cypher(
|
|
474
|
+
conn,
|
|
475
|
+
query,
|
|
476
|
+
{
|
|
477
|
+
"source": Jsonb(relation.source),
|
|
478
|
+
"target": Jsonb(relation.target),
|
|
479
|
+
},
|
|
480
|
+
)
|
|
481
|
+
|
|
482
|
+
await conn.commit()
|
|
483
|
+
|
|
484
|
+
logger.info(f"Successfully deleted {len(relations)} relations")
|
|
485
|
+
|
|
486
|
+
async def read_graph(self, limit: int = None) -> KnowledgeGraph:
|
|
487
|
+
"""Read the knowledge graph (up to ``limit`` entities)."""
|
|
488
|
+
return await self.load_graph(limit=limit)
|
|
489
|
+
|
|
490
|
+
async def search_memories(self, query: str, limit: int = None) -> KnowledgeGraph:
|
|
491
|
+
"""Search for memories based on a query (up to ``limit`` entities)."""
|
|
492
|
+
logger.info(f"Searching for memories with query: '{query}'")
|
|
493
|
+
return await self.load_graph(query, limit=limit)
|
|
494
|
+
|
|
495
|
+
async def find_memories_by_name(self, names: List[str]) -> KnowledgeGraph:
|
|
496
|
+
"""Find specific memories by their names."""
|
|
497
|
+
logger.info(f"Finding {len(names)} memories by name")
|
|
498
|
+
|
|
499
|
+
async with get_pool_connection(self.pool) as conn:
|
|
500
|
+
# Get entities
|
|
501
|
+
entity_query = """
|
|
502
|
+
MATCH (e:"Memory")
|
|
503
|
+
WHERE e.name IN %(names)s
|
|
504
|
+
RETURN e.name AS name, e.type AS type, e.observations AS observations
|
|
505
|
+
"""
|
|
506
|
+
|
|
507
|
+
entities_data = await self._execute_cypher(
|
|
508
|
+
conn, entity_query, {"names": Jsonb(names)}
|
|
509
|
+
)
|
|
510
|
+
|
|
511
|
+
entities = []
|
|
512
|
+
for record in entities_data:
|
|
513
|
+
entities.append(
|
|
514
|
+
Entity(
|
|
515
|
+
name=record.name,
|
|
516
|
+
type=record.type,
|
|
517
|
+
observations=record.observations or [],
|
|
518
|
+
)
|
|
519
|
+
)
|
|
520
|
+
|
|
521
|
+
# Get relations for found entities
|
|
522
|
+
relations = []
|
|
523
|
+
if entities:
|
|
524
|
+
rel_query = """
|
|
525
|
+
MATCH (source:"Memory")-[r]->(target:"Memory")
|
|
526
|
+
WHERE source.name IN %(names)s OR target.name IN %(names)s
|
|
527
|
+
RETURN source.name AS source, target.name AS target, label(r) AS "relationType"
|
|
528
|
+
"""
|
|
529
|
+
|
|
530
|
+
relations_data = await self._execute_cypher(
|
|
531
|
+
conn, rel_query, {"names": Jsonb(names)}
|
|
532
|
+
)
|
|
533
|
+
|
|
534
|
+
for record in relations_data:
|
|
535
|
+
relations.append(
|
|
536
|
+
Relation(
|
|
537
|
+
source=record.source,
|
|
538
|
+
target=record.target,
|
|
539
|
+
relationType=record.relationType,
|
|
540
|
+
)
|
|
541
|
+
)
|
|
542
|
+
|
|
543
|
+
await conn.commit()
|
|
544
|
+
|
|
545
|
+
logger.info(f"Found {len(entities)} entities and {len(relations)} relations")
|
|
546
|
+
return KnowledgeGraph(entities=entities, relations=relations)
|