Skip to content

Commit 49ce115

Browse files
author
Heiko
committed
Add on-prem graph memory backend
1 parent 96096ea commit 49ce115

16 files changed

Lines changed: 754 additions & 137 deletions

backend/app/config.py

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,15 @@ class Config:
3131
LLM_API_KEY = os.environ.get('LLM_API_KEY')
3232
LLM_BASE_URL = os.environ.get('LLM_BASE_URL', 'https://api.openai.com/v1')
3333
LLM_MODEL_NAME = os.environ.get('LLM_MODEL_NAME', 'gpt-4o-mini')
34+
35+
# Graph memory backend configuration
36+
GRAPH_MEMORY_BACKEND = os.environ.get('GRAPH_MEMORY_BACKEND', 'zep_cloud')
37+
GRAPHITI_MODEL_NAME = os.environ.get('GRAPHITI_MODEL_NAME', LLM_MODEL_NAME)
38+
GRAPHITI_EMBEDDING_MODEL_NAME = os.environ.get('GRAPHITI_EMBEDDING_MODEL_NAME', 'text-embedding-3-small')
39+
GRAPHITI_BRIDGE_URL = os.environ.get('GRAPHITI_BRIDGE_URL', 'http://graphiti-bridge:8008')
40+
FALKORDB_HOST = os.environ.get('FALKORDB_HOST', 'localhost')
41+
FALKORDB_PORT = int(os.environ.get('FALKORDB_PORT', '6379'))
42+
FALKORDB_DATABASE = os.environ.get('FALKORDB_DATABASE', 'mirofish')
3443

3544
# Zep配置
3645
ZEP_API_KEY = os.environ.get('ZEP_API_KEY')
@@ -69,7 +78,8 @@ def validate(cls) -> list[str]:
6978
errors: list[str] = []
7079
if not cls.LLM_API_KEY:
7180
errors.append("LLM_API_KEY 未配置")
72-
if not cls.ZEP_API_KEY:
81+
graph_backend = (cls.GRAPH_MEMORY_BACKEND or 'zep_cloud').lower()
82+
if graph_backend in {'zep', 'zep_cloud', 'zep-cloud'} and not cls.ZEP_API_KEY:
7383
errors.append("ZEP_API_KEY 未配置")
7484
return errors
7585

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,13 @@
1+
"""Graph memory backend adapters."""
2+
3+
from .base import GraphMemoryAdapter
4+
from .factory import create_graph_memory_adapter
5+
from .graphiti_bridge_adapter import GraphitiBridgeGraphMemoryAdapter
6+
from .zep_cloud_adapter import ZepCloudGraphMemoryAdapter
7+
8+
__all__ = [
9+
"GraphMemoryAdapter",
10+
"GraphitiBridgeGraphMemoryAdapter",
11+
"ZepCloudGraphMemoryAdapter",
12+
"create_graph_memory_adapter",
13+
]

backend/app/graph_memory/base.py

Lines changed: 66 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,66 @@
1+
"""Graph memory adapter contracts.
2+
3+
This module defines the narrow graph-memory surface Mirofish needs. Concrete
4+
backends can implement it without leaking vendor SDK details into services.
5+
"""
6+
7+
from __future__ import annotations
8+
9+
from abc import ABC, abstractmethod
10+
from typing import Any, Protocol
11+
12+
13+
class GraphMemoryAdapter(ABC):
14+
"""Backend-neutral graph memory interface used by Mirofish services."""
15+
16+
@abstractmethod
17+
def create_graph(self, graph_id: str, name: str, description: str) -> Any:
18+
"""Create a graph and return the backend response."""
19+
20+
@abstractmethod
21+
def set_ontology(self, graph_id: str, ontology: dict[str, Any]) -> Any:
22+
"""Apply ontology definitions for a graph."""
23+
24+
@abstractmethod
25+
def add_text_batch(self, graph_id: str, chunks: list[str]) -> Any:
26+
"""Add a batch of text episodes to a graph."""
27+
28+
@abstractmethod
29+
def add_text(self, graph_id: str, text: str) -> Any:
30+
"""Add a single text episode to a graph."""
31+
32+
@abstractmethod
33+
def get_episode(self, episode_uuid: str) -> Any:
34+
"""Return one episode by UUID."""
35+
36+
@abstractmethod
37+
def get_all_nodes(self, graph_id: str) -> list[Any]:
38+
"""Return all nodes for a graph."""
39+
40+
@abstractmethod
41+
def get_all_edges(self, graph_id: str) -> list[Any]:
42+
"""Return all edges for a graph."""
43+
44+
@abstractmethod
45+
def search(self, graph_id: str, query: str, limit: int = 10, scope: str = "edges", **kwargs: Any) -> Any:
46+
"""Search graph memory."""
47+
48+
@abstractmethod
49+
def get_node(self, node_uuid: str) -> Any:
50+
"""Return one node by UUID."""
51+
52+
@abstractmethod
53+
def get_node_edges(self, node_uuid: str) -> list[Any]:
54+
"""Return edges related to one node."""
55+
56+
@abstractmethod
57+
def delete_graph(self, graph_id: str) -> Any:
58+
"""Delete a graph."""
59+
60+
61+
class SupportsRawClient(Protocol):
62+
"""Compatibility escape hatch for legacy code not yet adapter-native."""
63+
64+
@property
65+
def raw_client(self) -> Any:
66+
"""Return the underlying SDK client."""
Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,23 @@
1+
"""Graph memory adapter factory."""
2+
3+
from __future__ import annotations
4+
5+
from typing import Optional
6+
7+
from ..config import Config
8+
from .base import GraphMemoryAdapter
9+
from .zep_cloud_adapter import ZepCloudGraphMemoryAdapter
10+
11+
12+
def create_graph_memory_adapter(api_key: Optional[str] = None, backend: Optional[str] = None) -> GraphMemoryAdapter:
13+
selected_backend = (backend or Config.GRAPH_MEMORY_BACKEND).strip().lower()
14+
15+
if selected_backend in {"zep", "zep_cloud", "zep-cloud"}:
16+
return ZepCloudGraphMemoryAdapter(api_key=api_key or Config.ZEP_API_KEY)
17+
18+
if selected_backend in {"graphiti", "graphiti_core", "graphiti-core", "graphiti_bridge", "graphiti-bridge"}:
19+
from .graphiti_bridge_adapter import GraphitiBridgeGraphMemoryAdapter
20+
21+
return GraphitiBridgeGraphMemoryAdapter(api_key=api_key or Config.LLM_API_KEY)
22+
23+
raise ValueError(f"Unsupported GRAPH_MEMORY_BACKEND: {selected_backend}")
Lines changed: 128 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,128 @@
1+
"""HTTP adapter for the on-premise Graphiti bridge service."""
2+
3+
from __future__ import annotations
4+
5+
import json
6+
from types import SimpleNamespace
7+
from typing import Any
8+
from urllib.error import HTTPError
9+
from urllib.parse import quote, urlencode
10+
from urllib.request import Request, urlopen
11+
12+
from .base import GraphMemoryAdapter
13+
from ..config import Config
14+
15+
16+
class GraphitiBridgeGraphMemoryAdapter(GraphMemoryAdapter):
17+
"""Graph memory adapter backed by the local Graphiti bridge service."""
18+
19+
def __init__(self, api_key: str | None = None, base_url: str | None = None):
20+
self.base_url = (base_url or Config.GRAPHITI_BRIDGE_URL).rstrip("/")
21+
self._node_graph_index: dict[str, str] = {}
22+
23+
@property
24+
def raw_client(self) -> None:
25+
return None
26+
27+
def create_graph(self, graph_id: str, name: str, description: str) -> Any:
28+
return self._to_namespace(self._request("POST", "/graphs", {"graph_id": graph_id, "name": name, "description": description}))
29+
30+
def set_ontology(self, graph_id: str, ontology: dict[str, Any]) -> Any:
31+
return self._request("POST", f"/graphs/{quote(graph_id)}/ontology", ontology)
32+
33+
def add_text_batch(self, graph_id: str, chunks: list[str]) -> list[Any]:
34+
data = self._request("POST", f"/graphs/{quote(graph_id)}/episodes", {"chunks": chunks})
35+
return [self._episode(item) for item in data.get("episodes", [])]
36+
37+
def add_text(self, graph_id: str, text: str) -> Any:
38+
data = self._request("POST", f"/graphs/{quote(graph_id)}/episodes", {"text": text})
39+
episodes = data.get("episodes", [])
40+
return self._episode(episodes[0]) if episodes else self._episode({"uuid": None, "processed": True})
41+
42+
def get_episode(self, episode_uuid: str) -> Any:
43+
return self._episode({"uuid": episode_uuid, "processed": True})
44+
45+
def get_all_nodes(self, graph_id: str) -> list[Any]:
46+
data = self._request("GET", f"/graphs/{quote(graph_id)}/nodes")
47+
nodes = [self._node(item) for item in data.get("nodes", [])]
48+
for node in nodes:
49+
self._node_graph_index[node.uuid_] = graph_id
50+
return nodes
51+
52+
def get_all_edges(self, graph_id: str) -> list[Any]:
53+
data = self._request("GET", f"/graphs/{quote(graph_id)}/edges")
54+
return [self._edge(item) for item in data.get("edges", [])]
55+
56+
def search(self, graph_id: str, query: str, limit: int = 10, scope: str = "edges", **kwargs: Any) -> Any:
57+
data = self._request("POST", f"/graphs/{quote(graph_id)}/search", {"query": query, "limit": limit, "scope": scope})
58+
nodes = [self._node(item) for item in data.get("nodes", [])]
59+
for node in nodes:
60+
self._node_graph_index[node.uuid_] = graph_id
61+
return SimpleNamespace(edges=[self._edge(item) for item in data.get("edges", [])], nodes=nodes)
62+
63+
def get_node(self, node_uuid: str) -> Any:
64+
graph_id = self._node_graph_index.get(node_uuid)
65+
if not graph_id:
66+
return None
67+
query = urlencode({"graph_id": graph_id})
68+
data = self._request("GET", f"/nodes/{quote(node_uuid)}?{query}")
69+
node = data.get("node")
70+
return self._node(node) if node else None
71+
72+
def get_node_edges(self, node_uuid: str) -> list[Any]:
73+
graph_id = self._node_graph_index.get(node_uuid)
74+
if not graph_id:
75+
return []
76+
query = urlencode({"graph_id": graph_id})
77+
data = self._request("GET", f"/nodes/{quote(node_uuid)}/edges?{query}")
78+
return [self._edge(item) for item in data.get("edges", [])]
79+
80+
def delete_graph(self, graph_id: str) -> Any:
81+
return self._request("DELETE", f"/graphs/{quote(graph_id)}")
82+
83+
def _request(self, method: str, path: str, payload: dict[str, Any] | None = None) -> dict[str, Any]:
84+
body = None if payload is None else json.dumps(payload).encode("utf-8")
85+
headers = {"Content-Type": "application/json"}
86+
req = Request(f"{self.base_url}{path}", data=body, headers=headers, method=method)
87+
try:
88+
with urlopen(req, timeout=120) as response:
89+
raw = response.read().decode("utf-8")
90+
return json.loads(raw) if raw else {}
91+
except HTTPError as exc:
92+
error_body = exc.read().decode("utf-8", errors="replace")
93+
raise RuntimeError(f"Graphiti bridge request failed: {exc.code} {error_body}") from exc
94+
95+
def _episode(self, data: dict[str, Any]) -> Any:
96+
uuid = data.get("uuid") or data.get("uuid_")
97+
return SimpleNamespace(uuid_=uuid, uuid=uuid, processed=data.get("processed", True))
98+
99+
def _node(self, data: dict[str, Any]) -> Any:
100+
uuid = data.get("uuid") or data.get("uuid_") or ""
101+
return SimpleNamespace(
102+
uuid_=uuid,
103+
uuid=uuid,
104+
name=data.get("name") or "",
105+
labels=data.get("labels") or [],
106+
summary=data.get("summary") or "",
107+
attributes=data.get("attributes") or {},
108+
created_at=data.get("created_at"),
109+
)
110+
111+
def _edge(self, data: dict[str, Any]) -> Any:
112+
uuid = data.get("uuid") or data.get("uuid_") or ""
113+
return SimpleNamespace(
114+
uuid_=uuid,
115+
uuid=uuid,
116+
name=data.get("name") or "",
117+
fact=data.get("fact") or "",
118+
source_node_uuid=data.get("source_node_uuid") or "",
119+
target_node_uuid=data.get("target_node_uuid") or "",
120+
attributes=data.get("attributes") or {},
121+
created_at=data.get("created_at"),
122+
valid_at=data.get("valid_at"),
123+
invalid_at=data.get("invalid_at"),
124+
expired_at=data.get("expired_at"),
125+
)
126+
127+
def _to_namespace(self, data: dict[str, Any]) -> Any:
128+
return SimpleNamespace(**data)
Lines changed: 121 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,121 @@
1+
"""Zep Cloud graph memory adapter."""
2+
3+
from __future__ import annotations
4+
5+
import warnings
6+
from typing import Any, Optional
7+
8+
from pydantic import Field
9+
from zep_cloud import EpisodeData, EntityEdgeSourceTarget
10+
from zep_cloud.client import Zep
11+
from zep_cloud.external_clients.ontology import EntityModel, EntityText, EdgeModel
12+
13+
from .base import GraphMemoryAdapter
14+
from ..utils.zep_paging import fetch_all_edges, fetch_all_nodes
15+
16+
17+
class ZepCloudGraphMemoryAdapter(GraphMemoryAdapter):
18+
"""Adapter preserving the existing Zep Cloud behavior."""
19+
20+
RESERVED_NAMES = {"uuid", "name", "group_id", "name_embedding", "summary", "created_at"}
21+
22+
def __init__(self, api_key: str):
23+
if not api_key:
24+
raise ValueError("ZEP_API_KEY 未配置")
25+
self._client = Zep(api_key=api_key)
26+
27+
@property
28+
def raw_client(self) -> Zep:
29+
return self._client
30+
31+
def create_graph(self, graph_id: str, name: str, description: str) -> Any:
32+
return self._client.graph.create(graph_id=graph_id, name=name, description=description)
33+
34+
def set_ontology(self, graph_id: str, ontology: dict[str, Any]) -> Any:
35+
warnings.filterwarnings("ignore", category=UserWarning, module="pydantic")
36+
37+
entity_types: dict[str, type[EntityModel]] = {}
38+
for entity_def in ontology.get("entity_types", []):
39+
name = entity_def["name"]
40+
description = entity_def.get("description", f"A {name} entity.")
41+
attrs: dict[str, Any] = {"__doc__": description}
42+
annotations: dict[str, Any] = {}
43+
44+
for attr_def in entity_def.get("attributes", []):
45+
attr_name = self._safe_attr_name(attr_def["name"])
46+
attr_desc = attr_def.get("description", attr_name)
47+
attrs[attr_name] = Field(description=attr_desc, default=None)
48+
annotations[attr_name] = Optional[EntityText]
49+
50+
attrs["__annotations__"] = annotations
51+
entity_class = type(name, (EntityModel,), attrs)
52+
entity_class.__doc__ = description
53+
entity_types[name] = entity_class
54+
55+
edge_definitions: dict[str, tuple[type[EdgeModel], list[EntityEdgeSourceTarget]]] = {}
56+
for edge_def in ontology.get("edge_types", []):
57+
name = edge_def["name"]
58+
description = edge_def.get("description", f"A {name} relationship.")
59+
attrs = {"__doc__": description}
60+
annotations = {}
61+
62+
for attr_def in edge_def.get("attributes", []):
63+
attr_name = self._safe_attr_name(attr_def["name"])
64+
attr_desc = attr_def.get("description", attr_name)
65+
attrs[attr_name] = Field(description=attr_desc, default=None)
66+
annotations[attr_name] = Optional[str]
67+
68+
attrs["__annotations__"] = annotations
69+
class_name = "".join(word.capitalize() for word in name.split("_"))
70+
edge_class = type(class_name, (EdgeModel,), attrs)
71+
edge_class.__doc__ = description
72+
73+
source_targets = [
74+
EntityEdgeSourceTarget(source=st.get("source", "Entity"), target=st.get("target", "Entity"))
75+
for st in edge_def.get("source_targets", [])
76+
]
77+
if source_targets:
78+
edge_definitions[name] = (edge_class, source_targets)
79+
80+
if not entity_types and not edge_definitions:
81+
return None
82+
83+
return self._client.graph.set_ontology(
84+
graph_ids=[graph_id],
85+
entities=entity_types if entity_types else None,
86+
edges=edge_definitions if edge_definitions else None,
87+
)
88+
89+
def add_text_batch(self, graph_id: str, chunks: list[str]) -> Any:
90+
episodes = [EpisodeData(data=chunk, type="text") for chunk in chunks]
91+
return self._client.graph.add_batch(graph_id=graph_id, episodes=episodes)
92+
93+
def add_text(self, graph_id: str, text: str) -> Any:
94+
return self._client.graph.add(graph_id=graph_id, type="text", data=text)
95+
96+
def get_episode(self, episode_uuid: str) -> Any:
97+
return self._client.graph.episode.get(uuid_=episode_uuid)
98+
99+
def get_all_nodes(self, graph_id: str) -> list[Any]:
100+
return fetch_all_nodes(self._client, graph_id)
101+
102+
def get_all_edges(self, graph_id: str) -> list[Any]:
103+
return fetch_all_edges(self._client, graph_id)
104+
105+
def search(self, graph_id: str, query: str, limit: int = 10, scope: str = "edges", **kwargs: Any) -> Any:
106+
return self._client.graph.search(graph_id=graph_id, query=query, limit=limit, scope=scope, **kwargs)
107+
108+
def get_node(self, node_uuid: str) -> Any:
109+
return self._client.graph.node.get(uuid_=node_uuid)
110+
111+
def get_node_edges(self, node_uuid: str) -> list[Any]:
112+
return self._client.graph.node.get_entity_edges(node_uuid=node_uuid)
113+
114+
def delete_graph(self, graph_id: str) -> Any:
115+
return self._client.graph.delete(graph_id=graph_id)
116+
117+
@classmethod
118+
def _safe_attr_name(cls, attr_name: str) -> str:
119+
if attr_name.lower() in cls.RESERVED_NAMES:
120+
return f"entity_{attr_name}"
121+
return attr_name

0 commit comments

Comments
 (0)