diff --git a/backend/app/services/graph_builder.py b/backend/app/services/graph_builder.py index a882377f..308499c1 100644 --- a/backend/app/services/graph_builder.py +++ b/backend/app/services/graph_builder.py @@ -16,6 +16,7 @@ from zep_cloud import EpisodeData, EntityEdgeSourceTarget from ..config import Config from ..models.task import TaskManager, TaskStatus from ..utils.zep_paging import fetch_all_nodes, fetch_all_edges +from ..utils.ontology import normalize_ontology_attribute from .text_processor import TextProcessor from ..utils.locale import t, get_locale, set_locale @@ -233,13 +234,11 @@ class GraphBuilderService: annotations = {} for attr_def in entity_def.get("attributes", []): - # LLM may return attributes as strings instead of dicts - if isinstance(attr_def, str): - attr_name = safe_attr_name(attr_def) - attr_desc = attr_def - else: - attr_name = safe_attr_name(attr_def["name"]) # 使用安全名称 - attr_desc = attr_def.get("description", attr_name) + normalized = normalize_ontology_attribute(attr_def) + if normalized is None: + continue + attr_name = safe_attr_name(normalized["name"]) # 使用安全名称 + attr_desc = normalized["description"] # Zep API 需要 Field 的 description,这是必需的 attrs[attr_name] = Field(description=attr_desc, default=None) annotations[attr_name] = Optional[EntityText] # 类型注解 @@ -262,13 +261,11 @@ class GraphBuilderService: annotations = {} for attr_def in edge_def.get("attributes", []): - # LLM may return attributes as strings instead of dicts - if isinstance(attr_def, str): - attr_name = safe_attr_name(attr_def) - attr_desc = attr_def - else: - attr_name = safe_attr_name(attr_def["name"]) # 使用安全名称 - attr_desc = attr_def.get("description", attr_name) + normalized = normalize_ontology_attribute(attr_def) + if normalized is None: + continue + attr_name = safe_attr_name(normalized["name"]) # 使用安全名称 + attr_desc = normalized["description"] # Zep API 需要 Field 的 description,这是必需的 attrs[attr_name] = Field(description=attr_desc, default=None) annotations[attr_name] = Optional[str] # 边属性用str类型 @@ -513,4 +510,3 @@ class GraphBuilderService: def delete_graph(self, graph_id: str): """删除图谱""" self.client.graph.delete(graph_id=graph_id) - diff --git a/backend/app/services/ontology_generator.py b/backend/app/services/ontology_generator.py index 0d8432c4..f57211af 100644 --- a/backend/app/services/ontology_generator.py +++ b/backend/app/services/ontology_generator.py @@ -10,6 +10,7 @@ from typing import Dict, Any, List, Optional from ..utils.llm_client import LLMClient from ..utils.locale import get_language_instruction from ..utils.file_parser import split_text_into_chunks +from ..utils.ontology import normalize_ontology_attribute logger = logging.getLogger(__name__) @@ -432,10 +433,11 @@ class OntologyGenerator: entity_name_map[original_name] = entity["name"] if "attributes" not in entity: entity["attributes"] = [] - # Normalize attributes: LLM may return strings instead of dicts + # Normalize LLM output and discard unusable attribute definitions. entity["attributes"] = [ - attr if isinstance(attr, dict) else {"name": attr, "type": "text", "description": attr} + normalized for attr in entity["attributes"] + if (normalized := normalize_ontology_attribute(attr)) is not None ] if "examples" not in entity: entity["examples"] = [] @@ -461,10 +463,11 @@ class OntologyGenerator: edge["source_targets"] = [] if "attributes" not in edge: edge["attributes"] = [] - # Normalize attributes: LLM may return strings instead of dicts + # Normalize LLM output and discard unusable attribute definitions. edge["attributes"] = [ - attr if isinstance(attr, dict) else {"name": attr, "type": "text", "description": attr} + normalized for attr in edge["attributes"] + if (normalized := normalize_ontology_attribute(attr)) is not None ] if len(edge.get("description", "")) > 100: edge["description"] = edge["description"][:97] + "..." diff --git a/backend/app/utils/ontology.py b/backend/app/utils/ontology.py new file mode 100644 index 00000000..5c5ba578 --- /dev/null +++ b/backend/app/utils/ontology.py @@ -0,0 +1,29 @@ +"""Helpers for validating LLM-generated ontology structures.""" + +from typing import Any, Dict, Optional + + +def normalize_ontology_attribute(attribute: Any) -> Optional[Dict[str, Any]]: + """Return a safe attribute definition, or ``None`` for unusable values.""" + + if isinstance(attribute, str): + if not attribute.strip(): + return None + return { + "name": attribute, + "type": "text", + "description": attribute, + } + + if not isinstance(attribute, dict): + return None + + name = attribute.get("name") + if not isinstance(name, str) or not name.strip(): + return None + + normalized = dict(attribute) + description = normalized.get("description") + if not isinstance(description, str) or not description: + normalized["description"] = name + return normalized diff --git a/backend/tests/test_ontology_attributes.py b/backend/tests/test_ontology_attributes.py new file mode 100644 index 00000000..87304fdd --- /dev/null +++ b/backend/tests/test_ontology_attributes.py @@ -0,0 +1,64 @@ +from app.services.ontology_generator import OntologyGenerator +from app.services.graph_builder import GraphBuilderService +from app.utils.ontology import normalize_ontology_attribute + + +def test_normalize_string_attribute(): + assert normalize_ontology_attribute("role") == { + "name": "role", + "type": "text", + "description": "role", + } + + +def test_preserve_valid_dictionary_attribute(): + original = {"name": "role", "type": "text", "description": "Public role"} + assert normalize_ontology_attribute(original) == original + assert normalize_ontology_attribute(original) is not original + + +def test_reject_unusable_attribute_shapes(): + for value in (None, 7, [], {}, {"name": None}, {"name": ""}, " "): + assert normalize_ontology_attribute(value) is None + + +def test_generator_normalizes_entity_and_edge_attributes(): + result = OntologyGenerator(llm_client=object())._validate_and_process({ + "entity_types": [{"name": "speaker", "attributes": ["role", None]}], + "edge_types": [{"name": "quotes", "attributes": ["source_url", {}]}], + }) + + assert result["entity_types"][0]["attributes"] == [{ + "name": "role", + "type": "text", + "description": "role", + }] + assert result["edge_types"][0]["attributes"] == [{ + "name": "source_url", + "type": "text", + "description": "source_url", + }] + + +def test_graph_builder_safety_net_accepts_strings_and_skips_invalid_values(): + 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": ["role", None, {"name": "summary"}], + }], + "edge_types": [], + }) + + speaker = captured["entities"]["Speaker"] + assert set(speaker.__annotations__) == {"role", "entity_summary"}