use std::collections::HashMap; use axum::{ Extension, extract::{Path, Query, State}, response::Json }; use serde::{Deserialize, Serialize}; use sqlx::SqlitePool; use utoipa::{IntoParams, ToSchema}; use crate::{AuthUser, api::error::ApiError}; // Type aliases for complex SQL row types #[allow(clippy::type_complexity)] type EntityRow = (i64, String, String, Option, Option, Option, i64); #[allow(clippy::type_complexity)] type NeighborRow = (i64, String, String, Option, Option, Option, i64, String); type MentionRow = (i64, String, i64, String); /// 实体信息 #[derive(Debug, Serialize, ToSchema)] pub struct EntityInfo { pub id: i64, pub name: String, pub entity_type: String, pub properties: HashMap, pub file_id: Option, pub kb_id: Option, pub created_at: i64, } /// 实体搜索参数 #[derive(Debug, Deserialize, IntoParams)] pub struct EntitySearchParams { pub q: Option, // 搜索关键词 pub entity_type: Option, // 实体类型筛选 pub kb_id: Option, // 知识库ID筛选 pub file_id: Option, // 文件ID筛选 pub limit: Option, // 返回数量限制 } /// 实体查询 #[utoipa::path( get, path = "/api/v1/knowledge/graph/entities", operation_id = "graph_search_entities", tag = "graph", params(EntitySearchParams), responses( (status = 200, description = "成功返回实体列表", body = Vec), (status = 400, description = "请求参数错误") ) )] pub async fn search_entities( Query(params): Query, State(pool): State, Extension(auth_user): Extension, ) -> Result>, ApiError> { let is_admin = auth_user.is_admin(); // Check KB permission if kb_id is specified if let Some(kb_id) = params.kb_id { let perm = crate::api::knowledge_base::get_kb_permission(&pool, kb_id, &auth_user.user_id, is_admin).await; if !crate::api::knowledge_base::meets_requirement(perm.as_deref(), "viewer") { return Err(ApiError::Forbidden("Permission denied.".to_string())); } } let mut sql = "SELECT id, name, entity_type, properties, file_id, kb_id, created_at FROM graph_nodes WHERE 1=1".to_string(); // 构建查询条件 if params.q.is_some() { sql.push_str(" AND name LIKE '%' || ? || '%'"); } if params.entity_type.is_some() { sql.push_str(" AND entity_type = ?"); } if params.kb_id.is_some() { sql.push_str(" AND kb_id = ?"); } if params.file_id.is_some() { sql.push_str(" AND file_id = ?"); } sql.push_str(" ORDER BY created_at DESC"); if let Some(limit) = params.limit { sql.push_str(&format!(" LIMIT {}", limit)); } else { sql.push_str(" LIMIT 100"); } // 执行查询 let mut query = sqlx::query_as::<_, (i64, String, String, Option, Option, Option, i64)>(&sql); if let Some(ref q) = params.q { query = query.bind(q); } if let Some(ref entity_type) = params.entity_type { query = query.bind(entity_type); } if let Some(kb_id) = params.kb_id { query = query.bind(kb_id); } if let Some(file_id) = params.file_id { query = query.bind(file_id); } let rows = query.fetch_all(&pool).await?; // Preload accessible KB ids for filtering let allowed_kb_ids: std::collections::HashSet = if is_admin { std::collections::HashSet::new() } else { crate::api::knowledge_base::get_user_viewable_kb_ids(&pool, &auth_user.user_id, false) .await .into_iter() .collect() }; let entities: Vec = rows .into_iter() .filter_map(|(id, name, entity_type, properties_json, file_id, kb_id, created_at)| { // Filter by KB permission if let Some(kid) = kb_id && !is_admin && !allowed_kb_ids.contains(&kid) { return None; } let properties: HashMap = if let Some(ref json) = properties_json { serde_json::from_str(json).unwrap_or_default() } else { HashMap::new() }; Some(EntityInfo { id, name, entity_type, properties, file_id, kb_id, created_at }) }) .collect(); Ok(Json(entities)) } /// 实体详情(包含邻居和提及) #[derive(Debug, Serialize, ToSchema)] pub struct EntityDetail { pub entity: EntityInfo, pub neighbors: Vec, pub mentions: Vec, } #[derive(Debug, Serialize, ToSchema)] pub struct NeighborInfo { pub entity: EntityInfo, pub relation_type: String, pub direction: String, // "outgoing" or "incoming" } #[derive(Debug, Serialize, ToSchema)] pub struct MentionInfo { pub slice_id: i64, pub context: String, pub file_id: i64, pub filename: String, } /// 获取实体详情 #[utoipa::path( get, path = "/api/v1/knowledge/graph/entities/{id}", operation_id = "graph_get_entity", tag = "graph", params( ("id" = i64, Path, description = "实体 ID") ), responses( (status = 200, description = "成功返回实体详情", body = EntityDetail), (status = 404, description = "实体不存在") ) )] pub async fn get_entity( Path(id): Path, State(pool): State, Extension(auth_user): Extension, ) -> Result, ApiError> { let is_admin = auth_user.is_admin(); // 查询实体基本信息 let entity_sql = "SELECT id, name, entity_type, properties, file_id, kb_id, created_at FROM graph_nodes WHERE id = ?"; let entity_row: Option = sqlx::query_as(entity_sql).bind(id).fetch_optional(&pool).await?; let entity_row = entity_row.ok_or_else(|| ApiError::Internal("Entity not found".to_string()))?; // Check permission on the entity's KB if let Some(kb_id) = entity_row.5 { let perm = crate::api::knowledge_base::get_kb_permission(&pool, kb_id, &auth_user.user_id, is_admin).await; if !crate::api::knowledge_base::meets_requirement(perm.as_deref(), "viewer") { return Err(ApiError::Forbidden("Permission denied.".to_string())); } } let properties: HashMap = if let Some(ref json) = entity_row.3 { serde_json::from_str(json).unwrap_or_default() } else { HashMap::new() }; let entity = EntityInfo { id: entity_row.0, name: entity_row.1, entity_type: entity_row.2, properties, file_id: entity_row.4, kb_id: entity_row.5, created_at: entity_row.6, }; // 出边、入边、提及三条查询相互独立,并发执行 let outgoing_sql = "SELECT n.id, n.name, n.entity_type, n.properties, n.file_id, n.kb_id, n.created_at, e.relation_type \ FROM graph_edges e \ INNER JOIN graph_nodes n ON e.target_node_id = n.id \ WHERE e.source_node_id = ? LIMIT 50"; let incoming_sql = "SELECT n.id, n.name, n.entity_type, n.properties, n.file_id, n.kb_id, n.created_at, e.relation_type \ FROM graph_edges e \ INNER JOIN graph_nodes n ON e.source_node_id = n.id \ WHERE e.target_node_id = ? LIMIT 50"; let mentions_sql = "SELECT m.slice_id, m.context, s.file_id, f.filename \ FROM entity_mentions m \ INNER JOIN slices s ON m.slice_id = s.id \ INNER JOIN files f ON s.file_id = f.id \ WHERE m.node_id = ? LIMIT 20"; let (outgoing_rows, incoming_rows, mentions_rows): (Vec, Vec, Vec) = tokio::try_join!( sqlx::query_as(outgoing_sql).bind(id).fetch_all(&pool), sqlx::query_as(incoming_sql).bind(id).fetch_all(&pool), sqlx::query_as(mentions_sql).bind(id).fetch_all(&pool), )?; let mut neighbors = Vec::new(); for row in outgoing_rows { let properties: HashMap = if let Some(ref json) = row.3 { serde_json::from_str(json).unwrap_or_default() } else { HashMap::new() }; neighbors.push(NeighborInfo { entity: EntityInfo { id: row.0, name: row.1, entity_type: row.2, properties, file_id: row.4, kb_id: row.5, created_at: row.6, }, relation_type: row.7, direction: "outgoing".to_string(), }); } for row in incoming_rows { let properties: HashMap = if let Some(ref json) = row.3 { serde_json::from_str(json).unwrap_or_default() } else { HashMap::new() }; neighbors.push(NeighborInfo { entity: EntityInfo { id: row.0, name: row.1, entity_type: row.2, properties, file_id: row.4, kb_id: row.5, created_at: row.6, }, relation_type: row.7, direction: "incoming".to_string(), }); } let mentions: Vec = mentions_rows .into_iter() .map(|(slice_id, context, file_id, filename)| MentionInfo { slice_id, context, file_id, filename }) .collect(); Ok(Json(EntityDetail { entity, neighbors, mentions })) } /// 图统计信息 #[derive(Debug, Serialize, ToSchema)] pub struct GraphStats { pub node_count: i64, pub edge_count: i64, pub entity_types: HashMap, pub relation_types: HashMap, } #[derive(Debug, Deserialize, IntoParams)] pub struct StatsParams { pub kb_id: Option, pub file_id: Option, } /// 获取图统计信息 #[utoipa::path( get, path = "/api/v1/knowledge/graph/stats", operation_id = "graph_get_graph_stats", tag = "graph", params(StatsParams), responses( (status = 200, description = "成功返回图统计信息", body = GraphStats), (status = 400, description = "请求参数错误") ) )] pub async fn get_graph_stats( Query(params): Query, State(pool): State, Extension(auth_user): Extension, ) -> Result, ApiError> { let is_admin = auth_user.is_admin(); // Check KB permission if kb_id is specified if let Some(kb_id) = params.kb_id { let perm = crate::api::knowledge_base::get_kb_permission(&pool, kb_id, &auth_user.user_id, is_admin).await; if !crate::api::knowledge_base::meets_requirement(perm.as_deref(), "viewer") { return Err(ApiError::Forbidden("Permission denied.".to_string())); } } // 构建 WHERE 子句 let where_clause = if let Some(file_id) = params.file_id { format!("WHERE file_id = {}", file_id) } else if let Some(kb_id) = params.kb_id { format!("WHERE kb_id = {}", kb_id) } else { "".to_string() }; // 节点数量 let node_count_sql = format!("SELECT COUNT(*) FROM graph_nodes {}", where_clause); let node_count: (i64,) = sqlx::query_as(&node_count_sql).fetch_one(&pool).await?; // 边数量 let edge_count_sql = if let Some(file_id) = params.file_id { format!("SELECT COUNT(*) FROM graph_edges WHERE file_id = {}", file_id) } else if let Some(kb_id) = params.kb_id { format!( "SELECT COUNT(*) FROM graph_edges e \ INNER JOIN graph_nodes n ON e.source_node_id = n.id \ WHERE n.kb_id = {}", kb_id ) } else { "SELECT COUNT(*) FROM graph_edges".to_string() }; let edge_count: (i64,) = sqlx::query_as(&edge_count_sql).fetch_one(&pool).await?; // 实体类型分布 let entity_types_sql = format!("SELECT entity_type, COUNT(*) FROM graph_nodes {} GROUP BY entity_type", where_clause); let entity_types_rows: Vec<(String, i64)> = sqlx::query_as(&entity_types_sql).fetch_all(&pool).await?; let entity_types: HashMap = entity_types_rows.into_iter().collect(); // 关系类型分布 let relation_types_sql = if let Some(file_id) = params.file_id { format!( "SELECT relation_type, COUNT(*) FROM graph_edges \ WHERE file_id = {} GROUP BY relation_type", file_id ) } else if let Some(kb_id) = params.kb_id { format!( "SELECT e.relation_type, COUNT(*) FROM graph_edges e \ INNER JOIN graph_nodes n ON e.source_node_id = n.id \ WHERE n.kb_id = {} GROUP BY e.relation_type", kb_id ) } else { "SELECT relation_type, COUNT(*) FROM graph_edges GROUP BY relation_type".to_string() }; let relation_types_rows: Vec<(String, i64)> = sqlx::query_as(&relation_types_sql).fetch_all(&pool).await?; let relation_types: HashMap = relation_types_rows.into_iter().collect(); Ok(Json(GraphStats { node_count: node_count.0, edge_count: edge_count.0, entity_types, relation_types })) }