use langchainrust::embeddings::{Embeddings, MockEmbeddings};
use langchainrust::retrieval::SimilarityRetriever;
use langchainrust::vector_stores::{InMemoryVectorStore, VectorStore};
use langchainrust::{AdaptiveRAG, Document, OpenAIChat, OpenAIConfig};
use std::sync::Arc;
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
let api_key = std::env::var("OPENAI_API_KEY").expect("请设置 OPENAI_API_KEY 环境变量");
let base_url = std::env::var("OPENAI_BASE_URL")
.unwrap_or_else(|_| "https://api.openai.com/v1".to_string());
let llm = OpenAIChat::new(OpenAIConfig {
api_key,
base_url,
model: "gpt-4o-mini".to_string(),
..Default::default()
});
let store = Arc::new(InMemoryVectorStore::new());
let embeddings = Arc::new(MockEmbeddings::new(3));
let docs = vec![
Document::new("Rust 是一门系统编程语言,注重安全和性能。"),
Document::new("Rust 的所有权系统避免了数据竞争和空指针。"),
Document::new("Tokio 是 Rust 最流行的异步运行时。"),
Document::new("async-std 是另一个 Rust 异步运行时,API 更接近标准库。"),
];
let doc_texts: Vec<&str> = docs.iter().map(|d| d.content.as_str()).collect();
let doc_embeddings = embeddings.embed_documents(&doc_texts).await?;
store.add_documents(docs, doc_embeddings).await?;
let retriever = SimilarityRetriever::new(store, embeddings);
let rag = AdaptiveRAG::new(llm, retriever)
.with_retrieve_k(4)
.with_multi_query_count(3);
let result = rag.invoke("你好,今天天气怎么样?").await?;
print_result("[NoRetrieval] 闲聊", &result);
let result = rag.invoke("Rust 的所有权系统是什么?").await?;
print_result("[SingleSearch] 具体问题", &result);
let result = rag.invoke("对比 Tokio 和 async-std 的调度模型").await?;
print_result("[MultiQuery] 复杂问题", &result);
Ok(())
}
fn print_result(label: &str, result: &langchainrust::AdaptiveRAGResult) {
println!("\n{}", label);
println!(" 决策: {}", result.decision);
println!(" 回答: {}", result.answer);
println!(" 来源文档数: {}", result.sources.len());
}