ForcePilot/backend/package/yuxi/knowledge/graphs/adapters/upload.py
Wenjie Zhang 410dd47c14 refactor: 将后端代码迁移至 backend 目录
- 将 server/, src/, scripts/, test/ 等目录移动到 backend/ 目录下
- 使用 git rename 保留文件历史记录
- 更新 docker-compose.yml 和 api.Dockerfile 配置

WIP: 项目结构重构进行中
2026-03-24 11:08:12 +08:00

123 lines
4.5 KiB
Python

from typing import TYPE_CHECKING, Any
from yuxi.knowledge.graphs.adapters.base import BaseNeo4jAdapter
from .base import GraphAdapter, GraphMetadata
if TYPE_CHECKING:
pass
class UploadGraphAdapter(GraphAdapter):
"""Upload图谱适配器 (Upload Graph Adapter)"""
def __init__(self, graph_db_instance=None, config: dict[str, Any] = None):
super().__init__(config)
# 延迟导入以避免循环导入
from yuxi.knowledge.graphs.upload_graph_service import UploadGraphService
# 初始化业务逻辑服务
if graph_db_instance and hasattr(graph_db_instance, "_service"):
# 如果传入的是新的包装类 GraphDatabase
self.service = graph_db_instance._service
else:
self.service = graph_db_instance if graph_db_instance else UploadGraphService()
if not self.service.is_running():
self.service.start()
# 初始化查询适配器(使用纯查询操作)
self._db = BaseNeo4jAdapter()
def _get_metadata(self) -> GraphMetadata:
"""获取 Upload 图谱元数据"""
return GraphMetadata(
graph_type="upload",
id_field="id",
name_field="name",
supports_embedding=True,
supports_threshold=True,
)
async def query_nodes(self, keyword: str, **kwargs) -> dict[str, Any]:
params = self._normalize_query_params(keyword, kwargs)
# 如果关键词是 "*" 或者为空,则执行采样查询
if not params["keyword"] or params["keyword"] == "*":
# 使用 BaseNeo4jAdapter 的连通子图查询
num = kwargs.get("max_nodes", 100)
raw_results = self._db._get_sample_nodes_with_connections(
num=num,
label_filter="Upload",
)
else:
# 否则执行关键词搜索(使用 service 的查询功能)
raw_results = self.service.query_node(
keyword=params["keyword"],
threshold=params.get("threshold", 0.9),
kgdb_name=params.get("kgdb_name", "neo4j"),
hops=params.get("hops", 2),
return_format="graph",
)
return self._format_results(raw_results)
def normalize_node(self, raw_node: Any) -> dict[str, Any]:
"""
raw_node expected format: {id: str, name: str, ...}
"""
node_id = raw_node.get("id")
name = raw_node.get("name")
return self._create_standard_node(
node_id=node_id,
name=name,
entity_type="entity",
labels=["Entity", "Upload"],
properties=raw_node,
source="upload",
)
def normalize_edge(self, raw_edge: Any) -> dict[str, Any]:
"""
raw_edge expected format: {type: str, source_id: str, target_id: str, ...}
"""
# Generate an ID if not present (Upload graph edges might not have explicit ID in simple dict return)
edge_id = raw_edge.get("id")
if not edge_id:
edge_id = f"{raw_edge.get('source_id')}_{raw_edge.get('type')}_{raw_edge.get('target_id')}"
return self._create_standard_edge(
edge_id=edge_id,
source_id=raw_edge.get("source_id"),
target_id=raw_edge.get("target_id"),
edge_type=raw_edge.get("type"),
properties=raw_edge,
)
async def get_labels(self) -> list[str]:
"""获取所有标签 - 使用 UploadGraphService"""
kgdb_name = self.config.get("kgdb_name", "neo4j")
info = self.service.get_graph_info(graph_name=kgdb_name)
return info.get("labels", []) if info else []
def _normalize_query_params(self, keyword: str, kwargs: dict) -> dict[str, Any]:
# Map max_depth to hops if present
hops = kwargs.get("hops")
if hops is None:
hops = kwargs.get("max_depth", 2)
return {
"keyword": keyword,
"threshold": kwargs.get("threshold", 0.9),
"kgdb_name": kwargs.get("kgdb_name", "neo4j"),
"hops": hops,
"filters": kwargs.get("filters", {}),
"context": kwargs.get("context", {}),
}
def _format_results(self, raw_results: dict[str, list]) -> dict[str, list]:
nodes = [self.normalize_node(n) for n in raw_results.get("nodes", [])]
edges = [self.normalize_edge(e) for e in raw_results.get("edges", [])]
return {"nodes": nodes, "edges": edges}