fix: enforce Zep ontology contract limits
This commit is contained in:
parent
9506e6f5e0
commit
d2f9e56eea
|
|
@ -16,7 +16,11 @@ from zep_cloud import EpisodeData, EntityEdgeSourceTarget
|
||||||
from ..config import Config
|
from ..config import Config
|
||||||
from ..models.task import TaskManager, TaskStatus
|
from ..models.task import TaskManager, TaskStatus
|
||||||
from ..utils.zep_paging import fetch_all_nodes, fetch_all_edges
|
from ..utils.zep_paging import fetch_all_nodes, fetch_all_edges
|
||||||
from ..utils.ontology import normalize_ontology_attribute
|
from ..utils.ontology import (
|
||||||
|
MAX_ONTOLOGY_TYPES,
|
||||||
|
RESERVED_ONTOLOGY_ATTRIBUTE_NAMES,
|
||||||
|
normalize_ontology_attributes,
|
||||||
|
)
|
||||||
from .text_processor import TextProcessor
|
from .text_processor import TextProcessor
|
||||||
from ..utils.locale import t, get_locale, set_locale
|
from ..utils.locale import t, get_locale, set_locale
|
||||||
|
|
||||||
|
|
@ -214,18 +218,15 @@ class GraphBuilderService:
|
||||||
# 这是 Zep SDK 要求的用法,警告来自动态类创建,可以安全忽略
|
# 这是 Zep SDK 要求的用法,警告来自动态类创建,可以安全忽略
|
||||||
warnings.filterwarnings('ignore', category=UserWarning, module='pydantic')
|
warnings.filterwarnings('ignore', category=UserWarning, module='pydantic')
|
||||||
|
|
||||||
# Zep 保留名称,不能作为属性名
|
|
||||||
RESERVED_NAMES = {'uuid', 'name', 'group_id', 'name_embedding', 'summary', 'created_at'}
|
|
||||||
|
|
||||||
def safe_attr_name(attr_name: str) -> str:
|
def safe_attr_name(attr_name: str) -> str:
|
||||||
"""将保留名称转换为安全名称"""
|
"""将保留名称转换为安全名称"""
|
||||||
if attr_name.lower() in RESERVED_NAMES:
|
if attr_name.lower() in RESERVED_ONTOLOGY_ATTRIBUTE_NAMES:
|
||||||
return f"entity_{attr_name}"
|
return f"entity_{attr_name}"
|
||||||
return attr_name
|
return attr_name
|
||||||
|
|
||||||
# 动态创建实体类型
|
# 动态创建实体类型
|
||||||
entity_types = {}
|
entity_types = {}
|
||||||
for entity_def in ontology.get("entity_types", []):
|
for entity_def in ontology.get("entity_types", [])[:MAX_ONTOLOGY_TYPES]:
|
||||||
name = entity_def["name"]
|
name = entity_def["name"]
|
||||||
description = entity_def.get("description", f"A {name} entity.")
|
description = entity_def.get("description", f"A {name} entity.")
|
||||||
|
|
||||||
|
|
@ -233,10 +234,9 @@ class GraphBuilderService:
|
||||||
attrs = {"__doc__": description}
|
attrs = {"__doc__": description}
|
||||||
annotations = {}
|
annotations = {}
|
||||||
|
|
||||||
for attr_def in entity_def.get("attributes", []):
|
for normalized in normalize_ontology_attributes(
|
||||||
normalized = normalize_ontology_attribute(attr_def)
|
entity_def.get("attributes", [])
|
||||||
if normalized is None:
|
):
|
||||||
continue
|
|
||||||
attr_name = safe_attr_name(normalized["name"]) # 使用安全名称
|
attr_name = safe_attr_name(normalized["name"]) # 使用安全名称
|
||||||
attr_desc = normalized["description"]
|
attr_desc = normalized["description"]
|
||||||
# Zep API 需要 Field 的 description,这是必需的
|
# Zep API 需要 Field 的 description,这是必需的
|
||||||
|
|
@ -252,7 +252,7 @@ class GraphBuilderService:
|
||||||
|
|
||||||
# 动态创建边类型
|
# 动态创建边类型
|
||||||
edge_definitions = {}
|
edge_definitions = {}
|
||||||
for edge_def in ontology.get("edge_types", []):
|
for edge_def in ontology.get("edge_types", [])[:MAX_ONTOLOGY_TYPES]:
|
||||||
name = edge_def["name"]
|
name = edge_def["name"]
|
||||||
description = edge_def.get("description", f"A {name} relationship.")
|
description = edge_def.get("description", f"A {name} relationship.")
|
||||||
|
|
||||||
|
|
@ -260,10 +260,9 @@ class GraphBuilderService:
|
||||||
attrs = {"__doc__": description}
|
attrs = {"__doc__": description}
|
||||||
annotations = {}
|
annotations = {}
|
||||||
|
|
||||||
for attr_def in edge_def.get("attributes", []):
|
for normalized in normalize_ontology_attributes(
|
||||||
normalized = normalize_ontology_attribute(attr_def)
|
edge_def.get("attributes", [])
|
||||||
if normalized is None:
|
):
|
||||||
continue
|
|
||||||
attr_name = safe_attr_name(normalized["name"]) # 使用安全名称
|
attr_name = safe_attr_name(normalized["name"]) # 使用安全名称
|
||||||
attr_desc = normalized["description"]
|
attr_desc = normalized["description"]
|
||||||
# Zep API 需要 Field 的 description,这是必需的
|
# Zep API 需要 Field 的 description,这是必需的
|
||||||
|
|
@ -294,7 +293,9 @@ class GraphBuilderService:
|
||||||
if entity_types or edge_definitions:
|
if entity_types or edge_definitions:
|
||||||
self.client.graph.set_ontology(
|
self.client.graph.set_ontology(
|
||||||
graph_ids=[graph_id],
|
graph_ids=[graph_id],
|
||||||
entities=entity_types if entity_types else None,
|
# zep-cloud 3.13.0 iterates entities.items(), so an edge-only
|
||||||
|
# ontology must pass an empty dictionary rather than None.
|
||||||
|
entities=entity_types,
|
||||||
edges=edge_definitions if edge_definitions else None,
|
edges=edge_definitions if edge_definitions else None,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -10,7 +10,10 @@ from typing import Dict, Any, List, Optional
|
||||||
from ..utils.llm_client import LLMClient
|
from ..utils.llm_client import LLMClient
|
||||||
from ..utils.locale import get_language_instruction
|
from ..utils.locale import get_language_instruction
|
||||||
from ..utils.file_parser import split_text_into_chunks
|
from ..utils.file_parser import split_text_into_chunks
|
||||||
from ..utils.ontology import normalize_ontology_attribute
|
from ..utils.ontology import (
|
||||||
|
MAX_ONTOLOGY_TYPES,
|
||||||
|
normalize_ontology_attributes,
|
||||||
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
@ -128,7 +131,7 @@ B. **具体类型(8个,根据文本内容设计)**:
|
||||||
### 3. 属性设计
|
### 3. 属性设计
|
||||||
|
|
||||||
- 每个实体类型1-3个关键属性
|
- 每个实体类型1-3个关键属性
|
||||||
- **注意**:属性名不能使用 `name`、`uuid`、`group_id`、`created_at`、`summary`(这些是系统保留字)
|
- **注意**:属性名不能使用 `name`、`uuid`、`group_id`、`graph_id`、`created_at`、`summary`(这些是系统保留字)
|
||||||
- 推荐使用:`full_name`, `title`, `role`, `position`, `location`, `description` 等
|
- 推荐使用:`full_name`, `title`, `role`, `position`, `location`, `description` 等
|
||||||
|
|
||||||
## 实体类型参考
|
## 实体类型参考
|
||||||
|
|
@ -268,7 +271,7 @@ class OntologyGenerator:
|
||||||
2. 最后2个必须是兜底类型:Person(个人兜底)和 Organization(组织兜底)
|
2. 最后2个必须是兜底类型:Person(个人兜底)和 Organization(组织兜底)
|
||||||
3. 前8个是根据文本内容设计的具体类型
|
3. 前8个是根据文本内容设计的具体类型
|
||||||
4. 所有实体类型必须是现实中可以发声的主体,不能是抽象概念
|
4. 所有实体类型必须是现实中可以发声的主体,不能是抽象概念
|
||||||
5. 属性名不能使用 name、uuid、group_id 等保留字,用 full_name、org_name 等替代
|
5. 属性名不能使用 name、uuid、group_id、graph_id 等保留字,用 full_name、org_name 等替代
|
||||||
"""
|
"""
|
||||||
|
|
||||||
return message
|
return message
|
||||||
|
|
@ -431,14 +434,11 @@ class OntologyGenerator:
|
||||||
if entity["name"] != original_name:
|
if entity["name"] != original_name:
|
||||||
logger.warning(f"Entity type name '{original_name}' auto-converted to '{entity['name']}'")
|
logger.warning(f"Entity type name '{original_name}' auto-converted to '{entity['name']}'")
|
||||||
entity_name_map[original_name] = entity["name"]
|
entity_name_map[original_name] = entity["name"]
|
||||||
if "attributes" not in entity:
|
# Normalize LLM output, enforce Zep field limits, and guarantee at
|
||||||
entity["attributes"] = []
|
# least one property for every custom ontology type.
|
||||||
# Normalize LLM output and discard unusable attribute definitions.
|
entity["attributes"] = normalize_ontology_attributes(
|
||||||
entity["attributes"] = [
|
entity.get("attributes", [])
|
||||||
normalized
|
)
|
||||||
for attr in entity["attributes"]
|
|
||||||
if (normalized := normalize_ontology_attribute(attr)) is not None
|
|
||||||
]
|
|
||||||
if "examples" not in entity:
|
if "examples" not in entity:
|
||||||
entity["examples"] = []
|
entity["examples"] = []
|
||||||
# 确保description不超过100字符
|
# 确保description不超过100字符
|
||||||
|
|
@ -461,20 +461,17 @@ class OntologyGenerator:
|
||||||
st["target"] = entity_name_map[st["target"]]
|
st["target"] = entity_name_map[st["target"]]
|
||||||
if "source_targets" not in edge:
|
if "source_targets" not in edge:
|
||||||
edge["source_targets"] = []
|
edge["source_targets"] = []
|
||||||
if "attributes" not in edge:
|
# Normalize LLM output, enforce Zep field limits, and guarantee at
|
||||||
edge["attributes"] = []
|
# least one property for every custom ontology type.
|
||||||
# Normalize LLM output and discard unusable attribute definitions.
|
edge["attributes"] = normalize_ontology_attributes(
|
||||||
edge["attributes"] = [
|
edge.get("attributes", [])
|
||||||
normalized
|
)
|
||||||
for attr in edge["attributes"]
|
|
||||||
if (normalized := normalize_ontology_attribute(attr)) is not None
|
|
||||||
]
|
|
||||||
if len(edge.get("description", "")) > 100:
|
if len(edge.get("description", "")) > 100:
|
||||||
edge["description"] = edge["description"][:97] + "..."
|
edge["description"] = edge["description"][:97] + "..."
|
||||||
|
|
||||||
# Zep API 限制:最多 10 个自定义实体类型,最多 10 个自定义边类型
|
# Zep API 限制:最多 10 个自定义实体类型,最多 10 个自定义边类型
|
||||||
MAX_ENTITY_TYPES = 10
|
MAX_ENTITY_TYPES = MAX_ONTOLOGY_TYPES
|
||||||
MAX_EDGE_TYPES = 10
|
MAX_EDGE_TYPES = MAX_ONTOLOGY_TYPES
|
||||||
|
|
||||||
# 去重:按 name 去重,保留首次出现的
|
# 去重:按 name 去重,保留首次出现的
|
||||||
seen_names = set()
|
seen_names = set()
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,25 @@
|
||||||
"""Helpers for validating LLM-generated ontology structures."""
|
"""Helpers for validating LLM-generated ontology structures."""
|
||||||
|
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
|
||||||
|
MAX_ONTOLOGY_TYPES = 10
|
||||||
|
MAX_ONTOLOGY_ATTRIBUTES = 10
|
||||||
|
RESERVED_ONTOLOGY_ATTRIBUTE_NAMES = frozenset({
|
||||||
|
"uuid",
|
||||||
|
"name",
|
||||||
|
"group_id",
|
||||||
|
"graph_id",
|
||||||
|
"name_embedding",
|
||||||
|
"summary",
|
||||||
|
"created_at",
|
||||||
|
})
|
||||||
|
|
||||||
|
_FALLBACK_ATTRIBUTE = {
|
||||||
|
"name": "details",
|
||||||
|
"type": "text",
|
||||||
|
"description": "Additional details about this ontology type.",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def normalize_ontology_attribute(attribute: Any) -> Optional[Dict[str, Any]]:
|
def normalize_ontology_attribute(attribute: Any) -> Optional[Dict[str, Any]]:
|
||||||
|
|
@ -27,3 +46,24 @@ def normalize_ontology_attribute(attribute: Any) -> Optional[Dict[str, Any]]:
|
||||||
if not isinstance(description, str) or not description:
|
if not isinstance(description, str) or not description:
|
||||||
normalized["description"] = name
|
normalized["description"] = name
|
||||||
return normalized
|
return normalized
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_ontology_attributes(attributes: Any) -> List[Dict[str, Any]]:
|
||||||
|
"""Return a non-empty Zep-compatible attribute list within service limits."""
|
||||||
|
|
||||||
|
if not isinstance(attributes, list):
|
||||||
|
attributes = []
|
||||||
|
|
||||||
|
normalized_attributes: List[Dict[str, Any]] = []
|
||||||
|
for attribute in attributes:
|
||||||
|
normalized = normalize_ontology_attribute(attribute)
|
||||||
|
if normalized is None:
|
||||||
|
continue
|
||||||
|
normalized_attributes.append(normalized)
|
||||||
|
if len(normalized_attributes) == MAX_ONTOLOGY_ATTRIBUTES:
|
||||||
|
break
|
||||||
|
|
||||||
|
if not normalized_attributes:
|
||||||
|
normalized_attributes.append(dict(_FALLBACK_ATTRIBUTE))
|
||||||
|
|
||||||
|
return normalized_attributes
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,14 @@
|
||||||
from app.services.ontology_generator import OntologyGenerator
|
from app.services.ontology_generator import OntologyGenerator
|
||||||
from app.services.graph_builder import GraphBuilderService
|
from app.services.graph_builder import GraphBuilderService
|
||||||
from app.utils.ontology import normalize_ontology_attribute
|
from app.utils.ontology import (
|
||||||
|
MAX_ONTOLOGY_ATTRIBUTES,
|
||||||
|
normalize_ontology_attribute,
|
||||||
|
normalize_ontology_attributes,
|
||||||
|
)
|
||||||
|
from zep_cloud.external_clients.ontology import (
|
||||||
|
edge_model_to_api_schema,
|
||||||
|
entity_model_to_api_schema,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_normalize_string_attribute():
|
def test_normalize_string_attribute():
|
||||||
|
|
@ -22,6 +30,22 @@ def test_reject_unusable_attribute_shapes():
|
||||||
assert normalize_ontology_attribute(value) is None
|
assert normalize_ontology_attribute(value) is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_attribute_list_is_non_empty_and_capped_for_zep():
|
||||||
|
assert normalize_ontology_attributes(None) == [{
|
||||||
|
"name": "details",
|
||||||
|
"type": "text",
|
||||||
|
"description": "Additional details about this ontology type.",
|
||||||
|
}]
|
||||||
|
|
||||||
|
attributes = [None] + [f"field_{index}" for index in range(12)]
|
||||||
|
normalized = normalize_ontology_attributes(attributes)
|
||||||
|
|
||||||
|
assert len(normalized) == MAX_ONTOLOGY_ATTRIBUTES
|
||||||
|
assert [attribute["name"] for attribute in normalized] == [
|
||||||
|
f"field_{index}" for index in range(MAX_ONTOLOGY_ATTRIBUTES)
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
def test_generator_normalizes_entity_and_edge_attributes():
|
def test_generator_normalizes_entity_and_edge_attributes():
|
||||||
result = OntologyGenerator(llm_client=object())._validate_and_process({
|
result = OntologyGenerator(llm_client=object())._validate_and_process({
|
||||||
"entity_types": [{"name": "speaker", "attributes": ["role", None]}],
|
"entity_types": [{"name": "speaker", "attributes": ["role", None]}],
|
||||||
|
|
@ -40,6 +64,16 @@ def test_generator_normalizes_entity_and_edge_attributes():
|
||||||
}]
|
}]
|
||||||
|
|
||||||
|
|
||||||
|
def test_generator_adds_a_property_to_empty_custom_types():
|
||||||
|
result = OntologyGenerator(llm_client=object())._validate_and_process({
|
||||||
|
"entity_types": [{"name": "speaker", "attributes": []}],
|
||||||
|
"edge_types": [{"name": "quotes", "attributes": []}],
|
||||||
|
})
|
||||||
|
|
||||||
|
assert result["entity_types"][0]["attributes"][0]["name"] == "details"
|
||||||
|
assert result["edge_types"][0]["attributes"][0]["name"] == "details"
|
||||||
|
|
||||||
|
|
||||||
def test_graph_builder_safety_net_accepts_strings_and_skips_invalid_values():
|
def test_graph_builder_safety_net_accepts_strings_and_skips_invalid_values():
|
||||||
captured = {}
|
captured = {}
|
||||||
|
|
||||||
|
|
@ -62,3 +96,75 @@ def test_graph_builder_safety_net_accepts_strings_and_skips_invalid_values():
|
||||||
|
|
||||||
speaker = captured["entities"]["Speaker"]
|
speaker = captured["entities"]["Speaker"]
|
||||||
assert set(speaker.__annotations__) == {"role", "entity_summary"}
|
assert set(speaker.__annotations__) == {"role", "entity_summary"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_graph_builder_emits_a_pinned_zep_sdk_compatible_schema():
|
||||||
|
captured = {}
|
||||||
|
|
||||||
|
class GraphApi:
|
||||||
|
def set_ontology(self, **kwargs):
|
||||||
|
captured.update(kwargs)
|
||||||
|
|
||||||
|
class Client:
|
||||||
|
graph = GraphApi()
|
||||||
|
|
||||||
|
builder = object.__new__(GraphBuilderService)
|
||||||
|
builder.client = Client()
|
||||||
|
builder.set_ontology("graph-id", {
|
||||||
|
"entity_types": [{
|
||||||
|
"name": "Speaker",
|
||||||
|
"attributes": ["graph_id"] + [
|
||||||
|
f"field_{index}" for index in range(10)
|
||||||
|
],
|
||||||
|
}],
|
||||||
|
"edge_types": [{
|
||||||
|
"name": "MENTIONS",
|
||||||
|
"attributes": [],
|
||||||
|
"source_targets": [{"source": "Speaker", "target": "Speaker"}],
|
||||||
|
}],
|
||||||
|
})
|
||||||
|
|
||||||
|
assert captured["graph_ids"] == ["graph-id"]
|
||||||
|
|
||||||
|
speaker = captured["entities"]["Speaker"]
|
||||||
|
entity_schema = entity_model_to_api_schema(speaker, "Speaker")
|
||||||
|
assert len(entity_schema["properties"]) == MAX_ONTOLOGY_ATTRIBUTES
|
||||||
|
assert entity_schema["properties"][0] == {
|
||||||
|
"name": "entity_graph_id",
|
||||||
|
"type": "Text",
|
||||||
|
"description": "graph_id",
|
||||||
|
}
|
||||||
|
|
||||||
|
mentions, source_targets = captured["edges"]["MENTIONS"]
|
||||||
|
edge_schema = edge_model_to_api_schema(mentions, "MENTIONS")
|
||||||
|
assert edge_schema["properties"] == [{
|
||||||
|
"name": "details",
|
||||||
|
"type": "Text",
|
||||||
|
"description": "Additional details about this ontology type.",
|
||||||
|
}]
|
||||||
|
assert source_targets[0].source == "Speaker"
|
||||||
|
assert source_targets[0].target == "Speaker"
|
||||||
|
|
||||||
|
|
||||||
|
def test_graph_builder_passes_an_empty_entity_mapping_for_edge_only_ontology():
|
||||||
|
captured = {}
|
||||||
|
|
||||||
|
class GraphApi:
|
||||||
|
def set_ontology(self, **kwargs):
|
||||||
|
captured.update(kwargs)
|
||||||
|
|
||||||
|
class Client:
|
||||||
|
graph = GraphApi()
|
||||||
|
|
||||||
|
builder = object.__new__(GraphBuilderService)
|
||||||
|
builder.client = Client()
|
||||||
|
builder.set_ontology("graph-id", {
|
||||||
|
"entity_types": [],
|
||||||
|
"edge_types": [{
|
||||||
|
"name": "RELATED_TO",
|
||||||
|
"attributes": ["reason"],
|
||||||
|
"source_targets": [{"source": "Entity", "target": "Entity"}],
|
||||||
|
}],
|
||||||
|
})
|
||||||
|
|
||||||
|
assert captured["entities"] == {}
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue