Improve LightRAG knowledge search locator alignment
This commit is contained in:
@@ -11,6 +11,7 @@ use axum::body::Body;
|
||||
use axum::http::{header, StatusCode};
|
||||
use axum::response::Response;
|
||||
use serde_json::{json, Value};
|
||||
use std::collections::HashSet;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::broadcast;
|
||||
use tracing::{info, warn};
|
||||
@@ -47,6 +48,105 @@ pub struct AcpRunBridge {
|
||||
event_tx: broadcast::Sender<SseEvent>,
|
||||
}
|
||||
|
||||
fn collect_citation_markdowns_from_value(value: &Value) -> Vec<Value> {
|
||||
fn add_citation(text: &str, seen: &mut HashSet<String>, out: &mut Vec<Value>) {
|
||||
let citation = text.trim();
|
||||
if !citation.is_empty() && seen.insert(citation.to_string()) {
|
||||
out.push(json!({ "citationMarkdown": citation }));
|
||||
}
|
||||
}
|
||||
|
||||
fn add_reference_citations(
|
||||
references: &[Value],
|
||||
seen: &mut HashSet<String>,
|
||||
out: &mut Vec<Value>,
|
||||
) -> bool {
|
||||
let has_precise = references.iter().any(|reference| {
|
||||
reference
|
||||
.get("citationMarkdown")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|value| !value.trim().is_empty())
|
||||
&& reference.get("locatorDegraded").and_then(Value::as_bool) != Some(true)
|
||||
});
|
||||
let mut added = false;
|
||||
for reference in references {
|
||||
if out.len() >= 8 {
|
||||
break;
|
||||
}
|
||||
let Some(citation) = reference.get("citationMarkdown").and_then(Value::as_str) else {
|
||||
continue;
|
||||
};
|
||||
if has_precise
|
||||
&& reference.get("locatorDegraded").and_then(Value::as_bool) == Some(true)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
let before = out.len();
|
||||
add_citation(citation, seen, out);
|
||||
added = added || out.len() > before;
|
||||
}
|
||||
added
|
||||
}
|
||||
|
||||
fn visit(value: &Value, seen: &mut HashSet<String>, out: &mut Vec<Value>) {
|
||||
if out.len() >= 8 {
|
||||
return;
|
||||
}
|
||||
match value {
|
||||
Value::String(text) => {
|
||||
let trimmed = text.trim();
|
||||
if (trimmed.starts_with('{') || trimmed.starts_with('['))
|
||||
&& trimmed.contains("citationMarkdown")
|
||||
{
|
||||
if let Ok(parsed) = serde_json::from_str::<Value>(trimmed) {
|
||||
visit(&parsed, seen, out);
|
||||
} else if let Some(first_line) = trimmed.lines().next() {
|
||||
if let Ok(parsed) = serde_json::from_str::<Value>(first_line.trim()) {
|
||||
visit(&parsed, seen, out);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Value::Array(items) => {
|
||||
for item in items {
|
||||
visit(item, seen, out);
|
||||
if out.len() >= 8 {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
Value::Object(map) => {
|
||||
let has_filtered_references = map
|
||||
.get("references")
|
||||
.and_then(Value::as_array)
|
||||
.is_some_and(|references| add_reference_citations(references, seen, out));
|
||||
if let Some(citation) = map.get("citationMarkdown").and_then(Value::as_str) {
|
||||
if !has_filtered_references {
|
||||
add_citation(citation, seen, out);
|
||||
}
|
||||
}
|
||||
for (key, item) in map {
|
||||
if has_filtered_references
|
||||
&& matches!(key.as_str(), "references" | "citations" | "uiCitations")
|
||||
{
|
||||
continue;
|
||||
}
|
||||
visit(item, seen, out);
|
||||
if out.len() >= 8 {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
let mut seen = HashSet::new();
|
||||
let mut out = Vec::new();
|
||||
visit(value, &mut seen, &mut out);
|
||||
out
|
||||
}
|
||||
|
||||
impl AcpRunBridge {
|
||||
/// Create a new ACP run: create session + start prompt in background.
|
||||
///
|
||||
@@ -213,6 +313,8 @@ pub fn acp_event_to_sse(event: AcpSessionEvent) -> Option<SseEvent> {
|
||||
status,
|
||||
content,
|
||||
} => {
|
||||
let output = json!(content);
|
||||
let citation_markdowns = collect_citation_markdowns_from_value(&output);
|
||||
let error = status == crate::acp_types::ToolCallStatus::Failed;
|
||||
let event = if error {
|
||||
"tool.failed"
|
||||
@@ -227,7 +329,8 @@ pub fn acp_event_to_sse(event: AcpSessionEvent) -> Option<SseEvent> {
|
||||
"toolCallId": tool_call_id,
|
||||
"status": status,
|
||||
"error": error,
|
||||
"output": content,
|
||||
"output": output,
|
||||
"citationMarkdowns": citation_markdowns,
|
||||
}),
|
||||
})
|
||||
}
|
||||
@@ -426,6 +529,48 @@ mod tests {
|
||||
assert_eq!(running.data["status"], "in_progress");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn acp_tool_completed_extracts_precise_ui_citations_from_prefixed_text() {
|
||||
let prefix = json!({
|
||||
"schema": "mnote.acp.tool_result_ui_citations.v1",
|
||||
"references": [{
|
||||
"citationMarkdown": "[来源定位降级:a.md](/documents/a)",
|
||||
"locatorDegraded": true
|
||||
}, {
|
||||
"citationMarkdown": "[b.md · p.2](/documents/b?page=2)",
|
||||
"locatorDegraded": false
|
||||
}],
|
||||
"uiCitations": [{
|
||||
"citationMarkdown": "[来源定位降级:a.md](/documents/a)"
|
||||
}]
|
||||
})
|
||||
.to_string();
|
||||
let completed = acp_event_to_sse(AcpSessionEvent::ToolCallUpdate {
|
||||
tool_call_id: "tool_2".into(),
|
||||
status: crate::acp_types::ToolCallStatus::Completed,
|
||||
content: Some(vec![crate::acp_types::ContentBlockWrapper {
|
||||
wrapper_type: "content".into(),
|
||||
content: crate::acp_types::TextContent {
|
||||
content_type: "text".into(),
|
||||
text: format!("{prefix}\n工具正文"),
|
||||
},
|
||||
}]),
|
||||
})
|
||||
.expect("tool complete");
|
||||
|
||||
assert_eq!(
|
||||
completed.data["citationMarkdowns"][0]["citationMarkdown"].as_str(),
|
||||
Some("[b.md · p.2](/documents/b?page=2)")
|
||||
);
|
||||
assert_eq!(
|
||||
completed.data["citationMarkdowns"]
|
||||
.as_array()
|
||||
.unwrap()
|
||||
.len(),
|
||||
1
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn acp_session_info_update_emits_session_info_updated_sse() {
|
||||
let sse = acp_event_to_sse(AcpSessionEvent::SessionInfoUpdate {
|
||||
|
||||
Reference in New Issue
Block a user