use std::sync::Mutex;
use async_trait::async_trait;
use futures::stream::StreamExt;
use pipecrab_core::{DataFrame, Decision, Direction, Finality, Processor, Role, SystemFrame, Transcript};
use pipecrab_runtime::{Outbound, Stage, StageError};
use crate::{Conversation, GenParams, LanguageModel, LmError, Message, TokenOut};
pub struct LmStage<M: LanguageModel> {
model: M,
params: GenParams,
convo: Mutex<Conversation>,
}
impl<M: LanguageModel> LmStage<M> {
pub fn new(model: M, system_prompt: impl Into<std::sync::Arc<str>>) -> Self {
Self::with_params(model, system_prompt, GenParams::default())
}
pub fn with_params(
model: M,
system_prompt: impl Into<std::sync::Arc<str>>,
params: GenParams,
) -> Self {
let convo = Conversation { messages: vec![Message::system(system_prompt)] };
Self { model, params, convo: Mutex::new(convo) }
}
}
pub struct Generate;
impl<M: LanguageModel> Processor for LmStage<M> {
type Effect = Generate;
fn decide_data(&mut self, frame: &DataFrame) -> Decision<Generate> {
match frame {
DataFrame::Transcript(Transcript {
role: Role::User,
finality: Finality::Final,
text,
}) => {
self.convo
.lock()
.expect("LmStage conversation mutex poisoned")
.messages
.push(Message::user(text.clone()));
Decision::drop().emit(Generate)
}
DataFrame::Transcript(Transcript { role: Role::User, finality: Finality::Partial { .. }, .. }) => {
Decision::drop()
}
_ => Decision::forward(),
}
}
fn decide_system(&mut self, _dir: Direction, frame: &SystemFrame) -> Decision<Generate> {
if matches!(frame, SystemFrame::Interrupt) {
self.model.cancel();
}
Decision::forward()
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl<M: LanguageModel> Stage for LmStage<M> {
async fn perform(&self, _effect: Generate, out: &Outbound) -> Result<(), StageError> {
let convo = { self.convo.lock().expect("LmStage conversation mutex poisoned").clone() };
let mut stream = self.model.generate(&convo, &self.params).await?;
let mut reply = String::new();
while let Some(item) = stream.next().await {
let TokenOut { delta } = item?;
reply.push_str(&delta);
let partial = Transcript::agent_partial(reply.clone());
debug_assert!(
matches!(partial.finality, Finality::Partial { stable } if stable == partial.text.len()),
"agent partial must be append-only (stable == text.len())",
);
let _ = out.send_data(partial.into()).await;
}
let _ = out.send_data(Transcript::agent_final(reply.clone()).into()).await;
self.convo
.lock()
.expect("LmStage conversation mutex poisoned")
.messages
.push(Message::assistant(reply));
Ok(())
}
}
impl From<LmError> for StageError {
fn from(e: LmError) -> Self {
StageError::new(e.to_string())
}
}