From da02a729a3584cfaf3e477b2b2ca1ddb066139e3 Mon Sep 17 00:00:00 2001 From: root Date: Thu, 30 Apr 2026 09:08:20 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E5=AE=9E=E4=BD=93=E9=BB=98=E8=AE=A4?= =?UTF-8?q?=E7=B1=BB=E5=9E=8BConcept=E8=80=8C=E9=9D=9ENULL=EF=BC=8C?= =?UTF-8?q?=E6=98=9F=E5=9B=BE=E5=9B=9E=E9=80=80=E6=97=B6=E9=97=B4=E6=88=B3?= =?UTF-8?q?=EF=BC=8Centity=5Ftypes=E6=94=B9=E4=B8=BAdict=E6=A0=BC=E5=BC=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- core/embedded_db.py | 5 ++-- core/graph_client.py | 14 +++++++---- core/tools/memory_tools.py | 10 +++++--- ui/static/graph.html | 50 -------------------------------------- 4 files changed, 18 insertions(+), 61 deletions(-) diff --git a/core/embedded_db.py b/core/embedded_db.py index f17637e..4965c38 100644 --- a/core/embedded_db.py +++ b/core/embedded_db.py @@ -393,8 +393,9 @@ class EmbeddedGraphDB: continue # 创建或更新实体 - for entity_name in [subject, obj]: - entity_type = entity_types.get(entity_name) if entity_types else None + for entity_name, entity_key in [(subject, 'subject_type'), (obj, 'object_type')]: + # 按优先级获取实体类型:1) triplet中的_type字段 2) entity_types字典 3) 默认 + entity_type = triplet.get(entity_key) or (entity_types.get(entity_name) if entity_types else None) or 'Concept' cursor.execute(""" INSERT INTO entities (name, type) diff --git a/core/graph_client.py b/core/graph_client.py index 046f424..4764cb2 100644 --- a/core/graph_client.py +++ b/core/graph_client.py @@ -121,7 +121,7 @@ class Neo4jGraph: return {"entities": list(entities.values()), "relations": relations[:20]} - def commit(self, triplets: list, entity_types: list = None, temporal_tag: str = None) -> dict: + def commit(self, triplets: list, entity_types: dict = None, temporal_tag: str = None) -> dict: """写入记忆""" global CURRENT_TURN with self.driver.session() as session: @@ -130,7 +130,6 @@ class Neo4jGraph: if not valid_triplets: return {"committed_count": 0, "details": []} - etype = entity_types[0] if entity_types else "unknown" date_bucket = temporal_tag or datetime.now().strftime("%Y-%m-%d") results = [] @@ -140,13 +139,17 @@ class Neo4jGraph: obj = triplet.get("object", "").strip() confidence = triplet.get("confidence", 0.9) + # 按优先级获取实体类型:1) triplet中的_type字段 2) entity_types字典 3) 默认 + s_type = triplet.get("subject_type") or (entity_types.get(subject) if entity_types else None) or "Concept" + o_type = triplet.get("object_type") or (entity_types.get(obj) if entity_types else None) or "Concept" + session.run(""" MERGE (s:Entity {name: $subject}) - ON CREATE SET s.type = $type, s.created_at = datetime(), s.mention_count = 1, s.updated_at = datetime() + ON CREATE SET s.type = $s_type, s.created_at = datetime(), s.mention_count = 1, s.updated_at = datetime() ON MATCH SET s.mention_count = coalesce(s.mention_count, 0) + 1, s.updated_at = datetime() MERGE (t:Entity {name: $object}) - ON CREATE SET t.type = $type, t.created_at = datetime(), t.mention_count = 1, t.updated_at = datetime() + ON CREATE SET t.type = $o_type, t.created_at = datetime(), t.mention_count = 1, t.updated_at = datetime() ON MATCH SET t.mention_count = coalesce(t.mention_count, 0) + 1, t.updated_at = datetime() CREATE (s)-[r:RELATES { @@ -159,7 +162,8 @@ class Neo4jGraph: confidence: $confidence, date_bucket: $date_bucket }]->(t) - """, subject=subject, object=obj, relation=relation, type=etype, + """, subject=subject, object=obj, relation=relation, + s_type=s_type, o_type=o_type, session_id=CURRENT_SESSION_ID, turn_id=CURRENT_TURN, confidence=confidence, date_bucket=date_bucket) diff --git a/core/tools/memory_tools.py b/core/tools/memory_tools.py index 9b1d1f3..4ce2bb4 100644 --- a/core/tools/memory_tools.py +++ b/core/tools/memory_tools.py @@ -103,16 +103,18 @@ MEMORY_TOOLS = [ "subject": {"type": "string"}, "relation": {"type": "string"}, "object": {"type": "string"}, - "confidence": {"type": "number"} + "confidence": {"type": "number"}, + "subject_type": {"type": "string", "description": "主体的实体类型,如 Person、Project"}, + "object_type": {"type": "string", "description": "客体的实体类型,如 Language、Technology"} }, "required": ["subject", "relation", "object"] }, "description": "三元组列表" }, "entity_types": { - "type": "array", - "items": {"type": "string"}, - "description": "实体类型(可选)" + "type": "object", + "additionalProperties": {"type": "string"}, + "description": "实体类型字典,如 {\"用户\": \"Person\", \"项目A\": \"Project\"}(可选)" }, "temporal_tag": { "type": "string", diff --git a/ui/static/graph.html b/ui/static/graph.html index 037941a..4caac28 100644 --- a/ui/static/graph.html +++ b/ui/static/graph.html @@ -548,7 +548,6 @@ let currentHighlightIds = []; let currentHighlightEdgeIds = new Set(); // 存储高亮的边ID let edgeParticles = {}; // 存储边的流动光点 {edgeId: {mesh, progress, edge}} - let edgeLabels = []; // 存储边标签 // 颜色映射 const typeColors = { @@ -813,13 +812,11 @@ function init() { // 清除旧的对象(包括边粒子) nodeMeshes.forEach(mesh => scene.remove(mesh)); edgeLines.forEach(line => scene.remove(line)); - edgeLabels.forEach(label => scene.remove(label)); Object.keys(edgeParticles).forEach(key => { scene.remove(edgeParticles[key].mesh); }); nodeMeshes = []; edgeLines = []; - edgeLabels = []; edgeParticles = {}; if (nodes.length === 0) return; @@ -1084,53 +1081,6 @@ function init() { scene.add(line); edgeLines.push(line); - - // 边标签(中间位置显示关系类型 + 时间戳) - const midX = (pos1.x + pos2.x) / 2; - const midY = (pos1.y + pos2.y) / 2; - const midZ = (pos1.z + pos2.z) / 2; - - const labelCanvas = document.createElement('canvas'); - const labelCtx = labelCanvas.getContext('2d'); - labelCanvas.width = 256; - labelCanvas.height = 48; - labelCtx.clearRect(0, 0, labelCanvas.width, labelCanvas.height); - - // 时间戳文本 - const ts = edge.created_at || ''; - const tsShort = ts ? ts.substring(0, 10) : ''; - let labelText = edge.relation_type; - if (tsShort) labelText += '\n' + tsShort; - - labelCtx.font = '16px Courier New'; - labelCtx.fillStyle = color; - labelCtx.textAlign = 'center'; - labelCtx.shadowColor = color; - labelCtx.shadowBlur = 4; - labelCtx.fillText(edge.relation_type, 128, 20); - if (tsShort) { - labelCtx.font = '12px Courier New'; - labelCtx.fillStyle = '#999999'; - labelCtx.fillText(tsShort, 128, 38); - } - - const labelTex = new THREE.CanvasTexture(labelCanvas); - labelTex.needsUpdate = true; - const labelMat = new THREE.SpriteMaterial({ - map: labelTex, - transparent: true, - opacity: 0.7, - depthTest: false, - depthWrite: false, - blending: THREE.AdditiveBlending - }); - const edgeLabel = new THREE.Sprite(labelMat); - edgeLabel.scale.set(5, 1, 1); - edgeLabel.position.set(midX, midY, midZ); - edgeLabel.userData = { edgeId: edge.id }; - - scene.add(edgeLabel); - edgeLabels.push(edgeLabel); }); // 调整相机位置以适应所有节点