fix: Neo4j DateTime serialization + match all edge types
- Convert Neo4j DateTime objects to strings via _safe_str() helper - Query all edge types (RELATES_TO + MENTIONS), not just RELATES_TO - Fix get_node_edges to match any relationship type Generated with [Claude Code](https://claude.ai/code) via [Happy](https://happy.engineering) Co-Authored-By: Claude <noreply@anthropic.com> Co-Authored-By: Happy <yesreply@happy.engineering>
This commit is contained in:
parent
c8eb8711cc
commit
339f2f2e50
|
|
@ -20,6 +20,13 @@ from ..utils.logger import get_logger
|
||||||
logger = get_logger('foresight.graphiti_client')
|
logger = get_logger('foresight.graphiti_client')
|
||||||
|
|
||||||
|
|
||||||
|
def _safe_str(val):
|
||||||
|
"""Convert any value to JSON-safe string, handling Neo4j DateTime etc."""
|
||||||
|
if val is None:
|
||||||
|
return None
|
||||||
|
return str(val)
|
||||||
|
|
||||||
|
|
||||||
from graphiti_core.embedder.client import EmbedderClient
|
from graphiti_core.embedder.client import EmbedderClient
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -354,8 +361,8 @@ class GraphitiClient:
|
||||||
name=n.get("name", ""),
|
name=n.get("name", ""),
|
||||||
labels=labels,
|
labels=labels,
|
||||||
summary=n.get("summary", ""),
|
summary=n.get("summary", ""),
|
||||||
attributes=dict(n) if n else {},
|
attributes={k: _safe_str(v) for k, v in dict(n).items()} if n else {},
|
||||||
created_at=str(n.get("created_at", "")) if n.get("created_at") else None,
|
created_at=_safe_str(n.get("created_at")),
|
||||||
))
|
))
|
||||||
return nodes
|
return nodes
|
||||||
|
|
||||||
|
|
@ -364,9 +371,9 @@ class GraphitiClient:
|
||||||
with self._driver.session() as session:
|
with self._driver.session() as session:
|
||||||
result = session.run(
|
result = session.run(
|
||||||
"""
|
"""
|
||||||
MATCH (a)-[r:RELATES_TO]->(b)
|
MATCH (a)-[r]->(b)
|
||||||
WHERE r.group_id = $gid
|
WHERE r.group_id = $gid
|
||||||
RETURN r, a.uuid as source_uuid, b.uuid as target_uuid
|
RETURN r, type(r) as rtype, a.uuid as source_uuid, b.uuid as target_uuid
|
||||||
""",
|
""",
|
||||||
gid=graph_id,
|
gid=graph_id,
|
||||||
)
|
)
|
||||||
|
|
@ -379,11 +386,11 @@ class GraphitiClient:
|
||||||
fact=r.get("fact", ""),
|
fact=r.get("fact", ""),
|
||||||
source_node_uuid=str(record["source_uuid"] or ""),
|
source_node_uuid=str(record["source_uuid"] or ""),
|
||||||
target_node_uuid=str(record["target_uuid"] or ""),
|
target_node_uuid=str(record["target_uuid"] or ""),
|
||||||
attributes=dict(r) if r else {},
|
attributes={k: _safe_str(v) for k, v in dict(r).items()} if r else {},
|
||||||
created_at=str(r.get("created_at", "")) if r.get("created_at") else None,
|
created_at=_safe_str(r.get("created_at")),
|
||||||
valid_at=str(r.get("valid_at", "")) if r.get("valid_at") else None,
|
valid_at=_safe_str(r.get("valid_at")),
|
||||||
invalid_at=str(r.get("invalid_at", "")) if r.get("invalid_at") else None,
|
invalid_at=_safe_str(r.get("invalid_at")),
|
||||||
expired_at=str(r.get("expired_at", "")) if r.get("expired_at") else None,
|
expired_at=_safe_str(r.get("expired_at")),
|
||||||
))
|
))
|
||||||
return edges
|
return edges
|
||||||
|
|
||||||
|
|
@ -407,8 +414,8 @@ class GraphitiClient:
|
||||||
name=n.get("name", ""),
|
name=n.get("name", ""),
|
||||||
labels=labels,
|
labels=labels,
|
||||||
summary=n.get("summary", ""),
|
summary=n.get("summary", ""),
|
||||||
attributes=dict(n) if n else {},
|
attributes={k: _safe_str(v) for k, v in dict(n).items()} if n else {},
|
||||||
created_at=str(n.get("created_at", "")) if n.get("created_at") else None,
|
created_at=_safe_str(n.get("created_at")),
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_node_edges(self, node_uuid: str) -> List[GraphitiEdge]:
|
def get_node_edges(self, node_uuid: str) -> List[GraphitiEdge]:
|
||||||
|
|
@ -416,7 +423,7 @@ class GraphitiClient:
|
||||||
with self._driver.session() as session:
|
with self._driver.session() as session:
|
||||||
result = session.run(
|
result = session.run(
|
||||||
"""
|
"""
|
||||||
MATCH (a)-[r:RELATES_TO]-(b)
|
MATCH (a)-[r]-(b)
|
||||||
WHERE a.uuid = $uuid
|
WHERE a.uuid = $uuid
|
||||||
RETURN r,
|
RETURN r,
|
||||||
CASE WHEN startNode(r) = a THEN a.uuid ELSE b.uuid END as source_uuid,
|
CASE WHEN startNode(r) = a THEN a.uuid ELSE b.uuid END as source_uuid,
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue