use async_trait::async_trait;
use futures_util::StreamExt;
use lc_callbacks::{RunTree, RunType};
use lc_core::runnables::RunnableConfig;
use lc_core::BaseChatModel;
use lc_providers::{wrap_chat_model, ProviderError};
use lc_schema::Message;
use serde_json::{json, Value};
use std::collections::HashMap;
use crate::base::{
stream_chain_with_callbacks, substitute_template, BaseChain, ChainError, ChainResult,
ChainStream, StreamToken,
};
use crate::BoxedChatModel;
pub struct LLMChain {
llm: BoxedChatModel,
prompt_template: String,
input_key: String,
output_key: String,
name: String,
}
impl LLMChain {
async fn stream_body(
&self,
inputs: HashMap<String, Value>,
config: Option<RunnableConfig>,
) -> Result<ChainStream, ChainError> {
self.validate_inputs(&inputs)?;
if config.as_ref().is_some_and(|c| c.is_cancelled()) {
return Err(ChainError::StreamError("Operation cancelled".to_string()));
}
let prompt = self.render_prompt(&inputs)?;
let messages = vec![Message::human(&prompt)];
let llm_stream = self
.llm
.stream_chat(messages, config)
.await
.map_err(|e| ChainError::StreamError(format!("LLM stream failed: {}", e)))?;
let stream = llm_stream.map(move |result| match result {
Ok(chunk) => Ok(StreamToken {
token: chunk.text,
is_final: false,
}),
Err(e) => Err(ChainError::StreamError(format!("Stream token error: {}", e))),
});
let final_stream = stream.chain(futures_util::stream::once(async move {
Ok(StreamToken {
token: String::new(),
is_final: true,
})
}));
Ok(Box::pin(final_stream))
}
pub fn new<L>(llm: L, prompt_template: impl Into<String>) -> Self
where
L: BaseChatModel + Send + Sync + 'static,
L::Error: Into<ProviderError>,
{
Self::from_wrapped(wrap_chat_model(llm), prompt_template)
}
pub(crate) fn from_wrapped(llm: BoxedChatModel, prompt_template: impl Into<String>) -> Self {
Self {
llm,
prompt_template: prompt_template.into(),
input_key: "question".to_string(),
output_key: "text".to_string(),
name: "llm_chain".to_string(),
}
}
pub fn with_input_key(mut self, key: impl Into<String>) -> Self {
self.input_key = key.into();
self
}
pub fn with_output_key(mut self, key: impl Into<String>) -> Self {
self.output_key = key.into();
self
}
pub fn with_name(mut self, name: impl Into<String>) -> Self {
self.name = name.into();
self
}
fn render_prompt(&self, inputs: &HashMap<String, Value>) -> Result<String, ChainError> {
let mut vars = HashMap::with_capacity(inputs.len());
for (key, value) in inputs {
let value_str = match value {
Value::String(s) => s.clone(),
_ => value.to_string(),
};
vars.insert(key.clone(), value_str);
}
let (prompt, missing) = substitute_template(&self.prompt_template, &vars);
if !missing.is_empty() {
return Err(ChainError::ExecutionError(format!(
"Prompt template has unreplaced variable(s): {}",
missing.join(", ")
)));
}
Ok(prompt)
}
}
#[async_trait]
impl BaseChain for LLMChain {
fn input_keys(&self) -> Vec<&str> {
vec![&self.input_key]
}
fn output_keys(&self) -> Vec<&str> {
vec![&self.output_key]
}
async fn invoke(&self, inputs: HashMap<String, Value>) -> Result<ChainResult, ChainError> {
self.validate_inputs(&inputs)?;
let prompt = self.render_prompt(&inputs)?;
let messages = vec![Message::human(&prompt)];
let result = self
.llm
.invoke(messages, None)
.await
.map_err(|e| ChainError::ExecutionError(format!("LLM call failed: {}", e)))?;
let mut output = HashMap::new();
output.insert(self.output_key.clone(), Value::String(result.content));
Ok(output)
}
async fn invoke_with_config(
&self,
inputs: HashMap<String, Value>,
config: Option<RunnableConfig>,
) -> Result<ChainResult, ChainError> {
self.validate_inputs(&inputs)?;
let callbacks = config.as_ref().and_then(|c| c.callbacks.clone());
let mut run = RunTree::new(self.name(), RunType::Chain, json!({ "inputs": inputs }));
if let Some(ref cb) = callbacks {
cb.dispatch_chain_start(&run, &run.inputs).await;
}
let prompt = match self.render_prompt(&inputs) {
Ok(p) => p,
Err(e) => {
let msg = e.to_string();
run.end_with_error(msg.clone());
if let Some(ref cb) = callbacks {
cb.dispatch_chain_error(&run, &msg).await;
}
return Err(e);
}
};
let messages = vec![Message::human(&prompt)];
let mut llm_run = run.create_child(
format!("{}.llm", self.name()),
RunType::Llm,
json!({"messages_count": messages.len()}),
);
if let Some(ref cb) = callbacks {
cb.dispatch_llm_start(&llm_run, &messages).await;
}
let llm_config = config.clone();
let result = self.llm.invoke(messages, llm_config).await;
match result {
Ok(llm_result) => {
llm_run.end(json!({"response": &llm_result.content}));
if let Some(ref cb) = callbacks {
cb.dispatch_llm_end(&llm_run, &llm_result.content).await;
}
let mut output = HashMap::new();
output.insert(
self.output_key.clone(),
Value::String(llm_result.content.clone()),
);
run.end(json!({"output": &llm_result.content}));
if let Some(ref cb) = callbacks {
cb.dispatch_chain_end(&run, &json!({"output": llm_result.content}))
.await;
}
Ok(output)
}
Err(e) => {
let err_msg = e.to_string();
llm_run.end_with_error(err_msg.clone());
if let Some(ref cb) = callbacks {
cb.dispatch_llm_error(&llm_run, &err_msg).await;
}
run.end_with_error(err_msg.clone());
if let Some(ref cb) = callbacks {
cb.dispatch_chain_error(&run, &err_msg).await;
}
Err(ChainError::ExecutionError(format!(
"LLM call failed: {}",
err_msg
)))
}
}
}
async fn stream(&self, inputs: HashMap<String, Value>) -> Result<ChainStream, ChainError> {
self.stream_body(inputs, None).await
}
async fn stream_with_config(
&self,
inputs: HashMap<String, Value>,
config: Option<RunnableConfig>,
) -> Result<ChainStream, ChainError> {
let output_key = Some(self.output_key.clone());
stream_chain_with_callbacks(
self.name(),
inputs,
config.clone(),
output_key,
|inputs| async move { self.stream_body(inputs, config).await },
)
.await
}
fn name(&self) -> &str {
&self.name
}
}
pub struct LLMChainBuilder {
llm: BoxedChatModel,
prompt_template: String,
input_key: Option<String>,
output_key: Option<String>,
name: Option<String>,
}
impl LLMChainBuilder {
pub fn new<L>(llm: L, prompt_template: impl Into<String>) -> Self
where
L: BaseChatModel + Send + Sync + 'static,
L::Error: Into<ProviderError>,
{
Self {
llm: wrap_chat_model(llm),
prompt_template: prompt_template.into(),
input_key: None,
output_key: None,
name: None,
}
}
pub fn input_key(mut self, key: impl Into<String>) -> Self {
self.input_key = Some(key.into());
self
}
pub fn output_key(mut self, key: impl Into<String>) -> Self {
self.output_key = Some(key.into());
self
}
pub fn name(mut self, name: impl Into<String>) -> Self {
self.name = Some(name.into());
self
}
pub fn build(self) -> LLMChain {
let mut chain = LLMChain::from_wrapped(self.llm, self.prompt_template);
if let Some(key) = self.input_key {
chain = chain.with_input_key(key);
}
if let Some(key) = self.output_key {
chain = chain.with_output_key(key);
}
if let Some(name) = self.name {
chain = chain.with_name(name);
}
chain
}
}