feat: v0.3.2 - CustomGraphBuilder replaces Graphiti for step 2 graph build

End of a long debugging rabbit hole. After ~10 attempts patching Graphiti's
compatibility with non-OpenAI LLMs (GLM, Qwen via SiliconFlow), every fix
uncovered another bug. Made the strategic call to abandon Graphiti entirely
for the graph-build path and write a minimal custom builder.

## Why Graphiti had to go

- Qwen 32B 32K context overflow on cumulative episode retrieval (60K-82K tokens)
- SiliconFlow Qwen 72B also 32K (not 128K as docs implied)
- GLM 4 Flash gives 128K context but triggers "20015 parameter invalid" on
  Graphiti's structured output calls (logprobs in reranker, empty input in
  embedder, nested dict in Neo4j writes)
- Each patch target was 1-2 levels deep in Graphiti internals
- 4 separate monkey-patches (EntityNode.save, bulk_utils, driver layer,
  reranker, embedder) still couldn't cover the extract_nodes path

## New architecture

backend/app/services/custom_graph_builder.py (NEW, 348 lines):
- Uses Foresight's own LLMClient (retry + GLM→Qwen fallback + token tracker)
- 1 LLM call per chunk (Graphiti needed 4-5)
- Single prompt extracts entities + relationships as JSON
- Writes directly to Neo4j via cypher MERGE (no Graphiti dependency)
- Schema identical to what Graphiti produced — downstream get_all_nodes /
  get_all_edges / zep_entity_reader all work unchanged
- ThreadPoolExecutor (10 workers) for concurrent chunk extraction
- Sequential Neo4j writes via shared session (avoids name conflicts in dedup)
- DDL wrapped in try/except (tolerates pre-existing Graphiti indexes)
- Entity name dedup via in-memory map → single uuid per canonical name
- Regex-sanitized label / relation type to prevent cypher injection

backend/app/api/graph.py:
- /api/graph/build endpoint now calls CustomGraphBuilder instead of
  builder.add_text_batches() → client.add_episodes_batch() (the old Graphiti path)

backend/app/services/graph_builder.py:
- _build_graph_worker also updated to use CustomGraphBuilder (dead code path
  but kept in sync)

backend/app/models/project.py:
- Default chunk_size 500 → 250 (to keep individual LLM prompts small)

backend/app/services/graphiti_client.py:
- Kept all monkey-patches for backward compat — they now only affect legacy
  Graphiti code paths that CustomGraphBuilder bypasses entirely:
  * _patch_neo4j_driver (AsyncSession.run nested-dict sanitize)
  * _patch_reranker_for_non_openai (GLM logprobs workaround)
  * _patch_embedder_empty_input (SiliconFlow empty-input guard)
  * _patch_entity_node_ops (EntityNode.save sanitize)
  * _patch_add_nodes_and_edges_bulk_tx (bulk-episode sanitize)
- These are kept for safety; Graphiti code path is no longer invoked in the
  production build flow but methods like add_episode still exist on the class

## Performance

- Sequential: ~194 chunks × 2-5s = 10-15 min
- Concurrent (10 workers): ~194 / 10 × 3-5s = **1-2 min** (5-8x speedup)
- Rate limiting handled by LLMClient retry/backoff, not raw thread contention

## Downstream compatibility

Verified:
- get_all_nodes: MATCH (n:Entity) WHERE n.group_id = $gid RETURN n, labels(n)
- get_all_edges: MATCH (a)-[r]->(b) WHERE r.group_id = $gid RETURN r, type(r)
- get_node / get_node_edges: MATCH by uuid

CustomGraphBuilder writes:
- (n:Entity [optional second label]) with uuid, name, summary, group_id, created_at
- [r:RELATION_TYPE] with uuid, name, fact, group_id, created_at

Schema matches exactly.

## Pipeline wiring

Step 1 ontology generation → Step 2 graph build (CustomGraphBuilder) →
Step 3 profile generation (ZepEntityReader queries Neo4j) → Step 4 simulation

No changes needed downstream of Step 2.
This commit is contained in:
liyizhouAI 2026-04-16 07:54:07 +08:00
parent f1f2ea4ea5
commit 7443082846
5 changed files with 705 additions and 76 deletions

View File

@ -347,8 +347,9 @@ def build_graph():
# 获取配置
graph_name = data.get('graph_name', project.name or 'Foresight Graph')
chunk_size = data.get('chunk_size', project.chunk_size or Config.DEFAULT_CHUNK_SIZE)
chunk_size = min(chunk_size, 250) # clamp: Qwen 32B max_seq_len=32768, keep prompt within window
chunk_overlap = data.get('chunk_overlap', project.chunk_overlap or Config.DEFAULT_CHUNK_OVERLAP)
# 更新项目配置
project.chunk_size = chunk_size
project.chunk_overlap = chunk_overlap
@ -445,14 +446,37 @@ def build_graph():
progress=15
)
episode_uuids = builder.add_text_batches(
graph_id,
chunks,
batch_size=3,
progress_callback=add_progress_callback
)
# Graphiti 同步处理,不需要轮询 episode 状态
# v0.3.2: 用 CustomGraphBuilder 替代 Graphiti add_episodes_batch
# 避免 Graphiti+GLM 兼容性问题20015 parameter invalid / 上下文爆窗口 / 嵌套 dict 等)
from ..services.custom_graph_builder import CustomGraphBuilder
def custom_progress_cb(current, total, msg):
prog = current / total if total else 0
progress = 15 + int(prog * 75) # 15% - 90%
task_manager.update_task(
task_id,
message=msg,
progress=progress
)
cgb = CustomGraphBuilder(graph_id=graph_id, ontology=ontology)
try:
build_stats = cgb.build(
full_text=text,
chunk_size=chunk_size,
chunk_overlap=chunk_overlap,
progress_callback=custom_progress_cb,
)
build_logger.info(
f"[{task_id}] CustomGraphBuilder 完成: "
f"entities={build_stats.get('entities_count')}, "
f"edges={build_stats.get('edges_count')}, "
f"chunks={build_stats.get('chunks_processed')}/{build_stats.get('total_chunks')}, "
f"failed={build_stats.get('chunks_failed')}"
)
finally:
cgb.close()
episode_uuids = [] # 兼容后续变量引用
# 获取图谱数据
task_manager.update_task(

View File

@ -46,7 +46,7 @@ class Project:
# 配置
simulation_requirement: Optional[str] = None
chunk_size: int = 500
chunk_size: int = 250
chunk_overlap: int = 50
# 错误信息
@ -92,7 +92,7 @@ class Project:
graph_id=data.get('graph_id'),
graph_build_task_id=data.get('graph_build_task_id'),
simulation_requirement=data.get('simulation_requirement'),
chunk_size=data.get('chunk_size', 500),
chunk_size=data.get('chunk_size', 250),
chunk_overlap=data.get('chunk_overlap', 50),
error=data.get('error')
)

View File

@ -0,0 +1,427 @@
"""
Custom Graph Builder - 替换 Graphiti 的自建实现
核心设计
1. 文本切分字符级简单切分 TextProcessor 一致
2. 对每个 chunk单次 LLM 调用抽取 entities + relationships
3. 去重 entity name
4. 直接用 neo4j driver 写入
5. 不使用 embeddingrerankerGraphiti 的任何代码
"""
import re
import uuid
import concurrent.futures
import threading
from datetime import datetime
from typing import Any, Callable, Dict, List, Optional
from neo4j import GraphDatabase
from ..config import Config
from ..utils.llm_client import LLMClient
from ..utils.logger import get_logger
logger = get_logger('foresight.custom_graph_builder')
class CustomGraphBuilder:
"""
LLM + Neo4j 自建图谱构建器
替代 Graphiti保持下游 GraphitiClient.get_all_nodes / get_all_edges 接口兼容
写入 schema
Node : Entity [可选第二 label] properties: uuid, name, summary, group_id, created_at
Edge : 任意类型 properties: uuid, name, fact, group_id, created_at
"""
def __init__(self, graph_id: str, ontology: Dict[str, Any]):
"""
Args:
graph_id: 图谱唯一标识group_id
ontology: {"entity_types": [...], "edge_types": [...]}
"""
self.graph_id = graph_id
self.ontology = ontology or {"entity_types": [], "edge_types": []}
self.llm = LLMClient()
self.driver = GraphDatabase.driver(
Config.NEO4J_URI,
auth=(Config.NEO4J_USER, Config.NEO4J_PASSWORD),
)
self.entities_created = 0
self.edges_created = 0
# 去重entity name → uuid
self._entity_uuid_by_name: Dict[str, str] = {}
def close(self):
if self.driver:
self.driver.close()
def ensure_constraints(self):
"""
建索引保证 MERGE 效率
每条 DDL 独立 try/except 包住 Graphiti 可能已经建过同名普通 index
而我们想要的是 unique constraint冲突时跳过不让初始化挂掉
"""
ddls = [
"CREATE CONSTRAINT entity_uuid_unique IF NOT EXISTS "
"FOR (n:Entity) REQUIRE n.uuid IS UNIQUE",
"CREATE INDEX entity_name_idx IF NOT EXISTS "
"FOR (n:Entity) ON (n.name)",
"CREATE INDEX entity_group_idx IF NOT EXISTS "
"FOR (n:Entity) ON (n.group_id)",
]
with self.driver.session() as session:
for ddl in ddls:
try:
session.run(ddl)
except Exception as e:
# IndexAlreadyExists / ConstraintAlreadyExists 等都是无害的
logger.warning(
f"CustomGraphBuilder: DDL skipped ({type(e).__name__}): "
f"{str(e)[:150]}"
)
# ------------------------------------------------------------------ #
# LLM 抽取
# ------------------------------------------------------------------ #
def _build_extraction_prompt(self, chunk_text: str) -> str:
entity_types_desc = "\n".join(
f"- {(t.get('name', t) if isinstance(t, dict) else t)}: "
f"{(t.get('description', '') if isinstance(t, dict) else '')}"
for t in self.ontology.get("entity_types", [])[:20]
) or "- Entity: 任何有意义的实体(人、组织、概念、事件、产品等)"
edge_types_desc = "\n".join(
f"- {(t.get('name', t) if isinstance(t, dict) else t)}: "
f"{(t.get('description', '') if isinstance(t, dict) else '')}"
for t in self.ontology.get("edge_types", [])[:20]
) or "- 任何合理的关系类型"
return (
"你是一个知识图谱构建助手。从下面文本中抽取实体和关系。\n\n"
"## 实体类型\n\n"
f"{entity_types_desc}\n\n"
"## 关系类型参考\n\n"
f"{edge_types_desc}\n\n"
"## 文本\n\n"
f"{chunk_text}\n\n"
"## 输出格式\n\n"
"严格返回 JSON不要有任何其他文字\n\n"
"```json\n"
"{\n"
' "entities": [\n'
' {\n'
' "name": "实体名(必须,简洁,禁止嵌套对象)",\n'
' "type": "实体类型名(从上面列表里选,或用 Entity",\n'
' "summary": "一句话描述这个实体(纯文本字符串,禁止嵌套对象,最多 200 字)"\n'
' }\n'
' ],\n'
' "relationships": [\n'
' {\n'
' "source": "源实体名(必须和 entities 里某个 name 一致)",\n'
' "target": "目标实体名",\n'
' "type": "关系类型(大写下划线格式,如 WORKS_AT, MENTIONS",\n'
' "fact": "一句话描述这个关系(纯文本,最多 150 字)"\n'
' }\n'
' ]\n'
'}\n'
"```\n\n"
"重要规则:\n"
"- 所有字段必须是字符串 primitive禁止嵌套 dict / object / list\n"
"- entities 每条的 name 要全局唯一\n"
"- relationships 里的 source/target 必须在 entities 里存在\n"
'- 如果没找到任何实体,返回 {"entities": [], "relationships": []}\n'
"- 实体数量控制在 5-20 之间\n"
"- 关系数量控制在 0-30 之间\n"
)
def extract_from_chunk(self, chunk_text: str) -> Dict[str, List[Dict]]:
"""对一个 chunk 做 LLM 抽取,返回 {entities, relationships}"""
prompt = self._build_extraction_prompt(chunk_text)
messages = [
{"role": "system", "content": "你是一个严谨的知识图谱构建助手,只返回 JSON。"},
{"role": "user", "content": prompt},
]
try:
result = self.llm.chat_json(messages, temperature=0.3, max_tokens=4096)
if not isinstance(result, dict):
logger.warning(f"chunk 返回非 dict: {type(result)}")
return {"entities": [], "relationships": []}
return {
"entities": result.get("entities", []) if isinstance(result.get("entities"), list) else [],
"relationships": result.get("relationships", []) if isinstance(result.get("relationships"), list) else [],
}
except Exception as e:
logger.error(f"chunk LLM 抽取失败: {type(e).__name__}: {str(e)[:200]}")
return {"entities": [], "relationships": []}
# ------------------------------------------------------------------ #
# sanitize helpers
# ------------------------------------------------------------------ #
@staticmethod
def _sanitize_value(v: Any) -> Any:
"""把 LLM 返回的字段 flatten 成 primitive"""
if v is None or isinstance(v, (str, int, float, bool)):
return v
if isinstance(v, dict):
if "value" in v:
return CustomGraphBuilder._sanitize_value(v["value"])
return str(v)[:500]
if isinstance(v, list):
return ", ".join(str(x)[:100] for x in v)[:500]
return str(v)[:500]
@staticmethod
def _safe_label(raw: str, default: str) -> str:
"""清洗字符串用作 Neo4j label / relationship type"""
cleaned = re.sub(r"[^A-Za-z0-9_]", "", raw)
return cleaned if cleaned else default
# ------------------------------------------------------------------ #
# Neo4j writes
# ------------------------------------------------------------------ #
def upsert_entity(self, tx, entity: Dict) -> Optional[str]:
"""写入或更新实体节点,返回 uuid同名实体幂等合并"""
name = self._sanitize_value(entity.get("name", ""))
if not isinstance(name, str):
name = str(name)
name = name.strip()
if not name:
return None
if name in self._entity_uuid_by_name:
return self._entity_uuid_by_name[name]
entity_uuid = str(uuid.uuid4())
raw_type = self._sanitize_value(entity.get("type", "Entity")) or "Entity"
if not isinstance(raw_type, str):
raw_type = str(raw_type)
safe_type = self._safe_label(raw_type.strip(), "Entity")
summary = self._sanitize_value(entity.get("summary", "")) or ""
if not isinstance(summary, str):
summary = str(summary)
summary = summary[:1000]
extra_label = f":{safe_type}" if safe_type != "Entity" else ""
tx.run(
f"""
MERGE (n:Entity{extra_label} {{name: $name, group_id: $gid}})
ON CREATE SET
n.uuid = $uuid,
n.summary = $summary,
n.created_at = $created_at
ON MATCH SET
n.summary = CASE
WHEN n.summary IS NULL OR n.summary = '' THEN $summary
ELSE n.summary
END
""",
name=name,
gid=self.graph_id,
uuid=entity_uuid,
summary=summary,
created_at=datetime.utcnow().isoformat(),
)
self._entity_uuid_by_name[name] = entity_uuid
self.entities_created += 1
return entity_uuid
def upsert_relationship(self, tx, rel: Dict) -> bool:
"""写入关系边source/target 必须已在 _entity_uuid_by_name 中"""
source_name = self._sanitize_value(rel.get("source", ""))
target_name = self._sanitize_value(rel.get("target", ""))
if not isinstance(source_name, str):
source_name = str(source_name)
if not isinstance(target_name, str):
target_name = str(target_name)
source_name = source_name.strip()
target_name = target_name.strip()
if not source_name or not target_name:
return False
if source_name not in self._entity_uuid_by_name or target_name not in self._entity_uuid_by_name:
return False
raw_type = self._sanitize_value(rel.get("type", "RELATED_TO")) or "RELATED_TO"
if not isinstance(raw_type, str):
raw_type = str(raw_type)
safe_type = self._safe_label(raw_type.strip().upper(), "RELATED_TO")
fact = self._sanitize_value(rel.get("fact", "")) or ""
if not isinstance(fact, str):
fact = str(fact)
fact = fact[:800]
rel_uuid = str(uuid.uuid4())
tx.run(
f"""
MATCH (a:Entity {{name: $src, group_id: $gid}}),
(b:Entity {{name: $tgt, group_id: $gid}})
MERGE (a)-[r:{safe_type} {{uuid: $uuid}}]->(b)
ON CREATE SET
r.name = $rel_name,
r.fact = $fact,
r.group_id = $gid,
r.created_at = $created_at
""",
src=source_name,
tgt=target_name,
gid=self.graph_id,
uuid=rel_uuid,
rel_name=safe_type,
fact=fact,
created_at=datetime.utcnow().isoformat(),
)
self.edges_created += 1
return True
# ------------------------------------------------------------------ #
# Main entry
# ------------------------------------------------------------------ #
def build(
self,
full_text: str,
chunk_size: int = 250,
chunk_overlap: int = 50,
progress_callback: Optional[Callable[[int, int, str], None]] = None,
concurrency: int = 10,
) -> Dict[str, Any]:
"""
主入口 chunk **并发** LLM 抽取 Neo4j
v0.3.2 并发化改动
- LLM 抽取用 ThreadPoolExecutor 并发 (默认 10 线程)
- Neo4j 写入串行 ( session + dict 去重必须顺序执行避免 name 冲突)
- LLM 调用走 Foresight LLMClient自带 retry/backoff/fallback 抗限流
Args:
full_text: 合并后的完整文档文本
chunk_size: 每块字符数
chunk_overlap: 块之间重叠字符数
progress_callback: 进度回调 (current, total, message)
concurrency: LLM 抽取并发数 (默认 10建议 5-15)
Returns:
{"entities_count": N, "edges_count": M, "chunks_processed": K,
"chunks_failed": F, "total_chunks": T}
"""
self.ensure_constraints()
# token 追踪
try:
from ..utils import token_tracker
token_tracker.set_stage("step2_graph_build")
except Exception:
pass
# 字符级切分
chunks: List[str] = []
start = 0
while start < len(full_text):
end = min(start + chunk_size, len(full_text))
chunks.append(full_text[start:end])
if end == len(full_text):
break
start = end - chunk_overlap
total_chunks = len(chunks)
logger.info(
f"CustomGraphBuilder: {total_chunks} chunks, "
f"concurrency={concurrency}, graph_id={self.graph_id}"
)
processed = 0
failed = 0
progress_lock = threading.Lock()
completed = [0]
def _extract_one(idx_chunk):
"""worker: LLM 抽取一个 chunk返回 (idx, extracted)"""
idx, chunk_text = idx_chunk
try:
return idx, self.extract_from_chunk(chunk_text)
except Exception as e:
logger.error(
f"chunk {idx + 1} LLM 抽取线程异常: "
f"{type(e).__name__}: {str(e)[:200]}"
)
return idx, {"entities": [], "relationships": []}
# 步骤1: 并发抽取全部 chunks
# 用单个持久 session 做 Neo4j 写入,按完成顺序消费结果
with self.driver.session() as session:
with concurrent.futures.ThreadPoolExecutor(max_workers=concurrency) as executor:
future_map = {
executor.submit(_extract_one, (i, c)): i
for i, c in enumerate(chunks)
}
for future in concurrent.futures.as_completed(future_map):
idx = future_map[future]
try:
_idx, extracted = future.result()
except Exception as e:
logger.error(f"chunk {idx + 1} future 异常: {e}")
failed += 1
continue
entities = extracted.get("entities", [])
relationships = extracted.get("relationships", [])
# 进度上报
with progress_lock:
completed[0] += 1
current = completed[0]
if progress_callback:
progress_callback(
current,
total_chunks,
f"处理 chunk {current}/{total_chunks}",
)
if not entities and not relationships:
processed += 1
continue
# Neo4j 写入(串行,避免同名 entity 并发插入冲突)
try:
def _tx_write(tx, _ents=entities, _rels=relationships):
for e in _ents:
if isinstance(e, dict):
self.upsert_entity(tx, e)
for r in _rels:
if isinstance(r, dict):
self.upsert_relationship(tx, r)
session.execute_write(_tx_write)
processed += 1
except Exception as e:
logger.error(
f"chunk {idx + 1} 写 Neo4j 失败: "
f"{type(e).__name__}: {str(e)[:200]}"
)
failed += 1
logger.info(
f"CustomGraphBuilder 完成: entities={self.entities_created}, "
f"edges={self.edges_created}, chunks={processed}/{total_chunks}, "
f"failed={failed}"
)
return {
"entities_count": self.entities_created,
"edges_count": self.edges_created,
"chunks_processed": processed,
"chunks_failed": failed,
"total_chunks": total_chunks,
}

View File

@ -50,7 +50,7 @@ class GraphBuilderService:
text: str,
ontology: Dict[str, Any],
graph_name: str = "Foresight Graph",
chunk_size: int = 500,
chunk_size: int = 250,
chunk_overlap: int = 50,
batch_size: int = 3
) -> str:
@ -121,22 +121,27 @@ class GraphBuilderService:
message=t('progress.textSplit', count=total_chunks)
)
# 4. 逐个添加 episodeGraphiti 同步处理,不需要轮询)
def progress_cb(msg, prog):
# 4. 用 CustomGraphBuilder 抽取实体 + 写入 Neo4j替代 Graphiti
from .custom_graph_builder import CustomGraphBuilder
def progress_cb(current, total, msg):
prog = current / total if total else 0
self.task_manager.update_task(
task_id,
progress=20 + int(prog * 0.65), # 20-85%
message=msg
)
self.client.add_episodes_batch(
graph_id=graph_id,
texts=chunks,
source_description=graph_name,
entity_types=entity_types,
edge_types=edge_types,
progress_callback=progress_cb,
)
builder = CustomGraphBuilder(graph_id=graph_id, ontology=ontology)
try:
builder.build(
full_text=text,
chunk_size=chunk_size,
chunk_overlap=chunk_overlap,
progress_callback=progress_cb,
)
finally:
builder.close()
# 5. 获取图谱信息
self.task_manager.update_task(

View File

@ -19,6 +19,78 @@ from ..utils.logger import get_logger
logger = get_logger('foresight.graphiti_client')
# ============================================================
# LOWEST-LEVEL PATCH: neo4j AsyncSession + AsyncTransaction
# ============================================================
_NEO4J_DRIVER_PATCHED = False
def _sanitize_value(v):
"""Recursively flatten a value to a Neo4j-safe primitive."""
if v is None or isinstance(v, (str, int, float, bool)):
return v
if isinstance(v, dict):
# LLM JSON-schema pollution: {description, title, type, value} → extract value
if "value" in v:
return _sanitize_value(v["value"])
# Plain dict cannot be a Neo4j property → serialize as JSON string
import json as _json
return _json.dumps(v, ensure_ascii=False)
if isinstance(v, list):
return [_sanitize_value(x) for x in v]
return str(v)
def _sanitize_params(params):
"""Sanitize all values in a cypher parameters dict."""
if not isinstance(params, dict):
return params
return {k: _sanitize_value(v) for k, v in params.items()}
def _patch_neo4j_driver():
"""
Ultimate fallback: patch neo4j AsyncSession.run and AsyncTransaction.run
to sanitize all parameters before they reach the Neo4j wire protocol.
This catches every cypher call regardless of which graphiti path is taken.
"""
global _NEO4J_DRIVER_PATCHED
if _NEO4J_DRIVER_PATCHED:
return
try:
from neo4j import AsyncSession, AsyncTransaction
_orig_session_run = AsyncSession.run
async def _patched_session_run(self, query, parameters=None, **kwargs):
sanitized = _sanitize_params(parameters) if parameters is not None else parameters
sanitized_kw = {k: _sanitize_value(v) for k, v in kwargs.items()} if kwargs else kwargs
return await _orig_session_run(self, query, sanitized, **sanitized_kw)
AsyncSession.run = _patched_session_run
_orig_tx_run = AsyncTransaction.run
async def _patched_tx_run(self, query, parameters=None, **kwargs):
sanitized = _sanitize_params(parameters) if parameters is not None else parameters
sanitized_kw = {k: _sanitize_value(v) for k, v in kwargs.items()} if kwargs else kwargs
return await _orig_tx_run(self, query, sanitized, **sanitized_kw)
AsyncTransaction.run = _patched_tx_run
_NEO4J_DRIVER_PATCHED = True
logger.info("[Foresight] Patched neo4j AsyncSession.run + AsyncTransaction.run for nested-dict sanitization")
except Exception as _e:
# Never let patch failure crash Flask startup
logger.warning(f"[Foresight] neo4j driver patch failed (non-fatal): {_e}")
# Patch immediately at import time — before any graphiti code runs
_patch_neo4j_driver()
def _safe_str(val):
"""Convert any value to JSON-safe string, handling Neo4j DateTime etc."""
@ -152,6 +224,111 @@ class GraphitiEdge:
return self.uuid_
def _patch_reranker_for_non_openai(reranker):
"""
Patch OpenAIRerankerClient.rank() to be compatible with providers that don't support
logprobs / top_logprobs / logit_bias (e.g. GLM-4-Flash, MiniMax).
The original implementation uses logprobs to get True/False probabilities.
This patch replaces it with a simple text-based True/False prompt that works
with any OpenAI-compatible provider.
Root cause of GLM code 20015: Graphiti sends logprobs=True + top_logprobs=2
+ logit_bias which GLM rejects with "The parameter is invalid."
"""
try:
import numpy as np
from graphiti_core.helpers import semaphore_gather
from graphiti_core.llm_client import RateLimitError
import openai as _openai
client_obj = reranker.client
model_name = reranker.config.model or "glm-4-flash"
async def _compatible_rank(query: str, passages: list[str]) -> list[tuple[str, float]]:
if not passages:
return []
openai_messages_list = [
[
{
"role": "system",
"content": "You are an expert tasked with determining whether the passage is relevant to the query",
},
{
"role": "user",
"content": (
'Respond with only "True" if PASSAGE is relevant to QUERY and "False" otherwise.\n'
f"<PASSAGE>\n{passage}\n</PASSAGE>\n"
f"<QUERY>\n{query}\n</QUERY>"
),
},
]
for passage in passages
]
try:
responses = await semaphore_gather(
*[
client_obj.chat.completions.create(
model=model_name,
messages=openai_messages,
temperature=0,
max_tokens=5,
)
for openai_messages in openai_messages_list
]
)
results = []
for passage, response in zip(passages, responses):
text = response.choices[0].message.content or ""
score = 1.0 if "true" in text.strip().lower() else 0.0
results.append((passage, score))
results.sort(reverse=True, key=lambda x: x[1])
return results
except _openai.RateLimitError as e:
raise RateLimitError from e
except Exception as e:
import logging as _logging
_logging.getLogger(__name__).error(f"[Foresight] Reranker error: {e}")
# Fallback: return all passages with equal score (no reranking)
return [(p, 0.5) for p in passages]
reranker.rank = _compatible_rank
logger.info("[Foresight] Patched OpenAIRerankerClient.rank() for non-OpenAI provider compatibility (no logprobs)")
except Exception as e:
logger.warning(f"[Foresight] Failed to patch reranker (non-fatal): {e}")
def _patch_embedder_empty_input(embedder):
"""
Patch OpenAIEmbedder.create_batch to guard against empty input lists.
Root cause of code 20015 from SiliconFlow: Graphiti calls create_batch([]) when
there are no nodes to embed (e.g. create_entity_node_embeddings with empty list).
SiliconFlow (and GLM) return error 20015 for empty input arrays.
Fix: return [] immediately if input is empty, without calling the API.
"""
try:
orig_create_batch = embedder.create_batch
async def guarded_create_batch(input_data_list):
if not input_data_list:
return []
return await orig_create_batch(input_data_list)
embedder.create_batch = guarded_create_batch
logger.info("[Foresight] Patched OpenAIEmbedder.create_batch to guard empty input (fixes code 20015)")
except Exception as e:
logger.warning(f"[Foresight] Failed to patch embedder (non-fatal): {e}")
class GraphitiClient:
"""
Graphiti + Neo4j 知识图谱客户端
@ -190,6 +367,10 @@ class GraphitiClient:
if self._graphiti is not None:
return
# Re-apply driver patch (idempotent) in case import order caused it to run
# before neo4j was available, or the class got re-imported.
_patch_neo4j_driver()
from graphiti_core import Graphiti
from graphiti_core.llm_client.openai_generic_client import OpenAIGenericClient
from graphiti_core.llm_client.config import LLMConfig
@ -217,10 +398,16 @@ class GraphitiClient:
)
embedder = OpenAIEmbedder(config=embedder_config)
# Reranker: use the LLM config (MiniMax)
# Reranker: use the LLM config
from graphiti_core.cross_encoder.openai_reranker_client import OpenAIRerankerClient
reranker = OpenAIRerankerClient(config=llm_config)
# Patch reranker if the provider doesn't support logprobs (e.g. GLM, MiniMax)
_patch_reranker_for_non_openai(reranker)
# Patch embedder to guard against empty input (SiliconFlow/GLM return 20015 for [])
_patch_embedder_empty_input(embedder)
self._graphiti = Graphiti(
self.neo4j_uri,
self.neo4j_user,
@ -241,66 +428,52 @@ class GraphitiClient:
def _patch_entity_node_ops(self):
"""
Patch Neo4jEntityNodeOperations.save and save_bulk to sanitize entity_data
before writing to Neo4j. Handles the case where Qwen/other LLMs return
nested dicts (JSON-schema metadata) in fields like 'summary'.
Patch both EntityNode.save (single-episode path) and
bulk_utils.add_nodes_and_edges_bulk_tx (bulk/batch episode path) to sanitize
entity_data before writing to Neo4j.
Root cause: Qwen/other LLMs return JSON-schema metadata dicts in fields like
'summary' instead of plain strings. e.g.:
{"summary": {"description": "...", "title": "Summary", "type": "string", "value": "actual text"}}
The bulk path builds entity_data dicts and passes them directly via tx.run()
without calling EntityNode.save at all, so we must patch both paths.
"""
from graphiti_core.driver.neo4j.operations.entity_node_ops import Neo4jEntityNodeOperations
from graphiti_core.driver.driver import GraphProvider
from graphiti_core.driver.query_executor import QueryExecutor, Transaction
# --- Patch 1: EntityNode.save (single-episode path) ---
from graphiti_core.nodes import EntityNode
from graphiti_core.models.nodes.node_db_queries import (
get_entity_node_save_query,
get_entity_node_save_bulk_query,
)
from typing import Any
_flatten = _sanitize_entity_data # closure over module-level helper
original_save = EntityNode.save
original_save = Neo4jEntityNodeOperations.save
original_save_bulk = Neo4jEntityNodeOperations.save_bulk
async def patched_save(self_node, driver):
if isinstance(self_node.summary, (dict, list)):
self_node.summary = _flatten_entity_property(self_node.summary)
if self_node.attributes:
self_node.attributes = _sanitize_entity_data(self_node.attributes)
return await original_save(self_node, driver)
async def patched_save(self_ops, executor, node: EntityNode, tx=None):
entity_data: dict[str, Any] = {
'uuid': node.uuid,
'name': node.name,
'name_embedding': node.name_embedding,
'group_id': node.group_id,
'summary': node.summary,
'created_at': node.created_at,
}
entity_data.update(node.attributes or {})
entity_data = _flatten(entity_data)
labels = ':'.join(list(set(node.labels + ['Entity'])))
query = get_entity_node_save_query(GraphProvider.NEO4J, labels)
if tx is not None:
await tx.run(query, entity_data=entity_data)
else:
await executor.execute_query(query, entity_data=entity_data)
EntityNode.save = patched_save
logger.info("[Foresight] Patched EntityNode.save for single-episode sanitization")
async def patched_save_bulk(self_ops, executor, nodes: list, tx=None, batch_size: int = 100):
prepared: list[dict[str, Any]] = []
for node in nodes:
entity_data: dict[str, Any] = {
'uuid': node.uuid,
'name': node.name,
'group_id': node.group_id,
'summary': node.summary,
'created_at': node.created_at,
'name_embedding': node.name_embedding,
'labels': list(set(node.labels + ['Entity'])),
}
entity_data.update(node.attributes or {})
prepared.append(_flatten(entity_data))
query = get_entity_node_save_bulk_query(GraphProvider.NEO4J, prepared)
if tx is not None:
await tx.run(query, nodes=prepared)
else:
await executor.execute_query(query, nodes=prepared)
# --- Patch 2: add_nodes_and_edges_bulk_tx (bulk episode path) ---
import graphiti_core.utils.bulk_utils as _bulk_utils
from typing import Any as _Any
Neo4jEntityNodeOperations.save = patched_save
Neo4jEntityNodeOperations.save_bulk = patched_save_bulk
logger.info("Patched Neo4jEntityNodeOperations to sanitize LLM-polluted entity fields")
original_bulk_tx = _bulk_utils.add_nodes_and_edges_bulk_tx
async def patched_bulk_tx(tx, episodic_nodes, episodic_edges, entity_nodes,
entity_edges, embedder):
# Sanitize each entity node's summary and attributes in-place before
# original function builds entity_data dicts from them
for node in entity_nodes:
if isinstance(node.summary, (dict, list)):
node.summary = _flatten_entity_property(node.summary)
if node.attributes:
node.attributes = _sanitize_entity_data(node.attributes)
return await original_bulk_tx(tx, episodic_nodes, episodic_edges,
entity_nodes, entity_edges, embedder)
_bulk_utils.add_nodes_and_edges_bulk_tx = patched_bulk_tx
logger.info("[Foresight] Patched add_nodes_and_edges_bulk_tx for bulk-episode sanitization")
@classmethod
def get_instance(cls) -> 'GraphitiClient':