use async_trait::async_trait;
use crate::error::{Error, Result};
use crate::tools::ToolDefinition;
use crate::types::Message;
#[derive(Debug, Clone, Default)]
pub struct SessionContext {
pub session_id: Option<String>,
pub service_session_id: Option<String>,
pub input_messages: Vec<Message>,
pub instructions: Option<String>,
pub messages: Vec<Message>,
pub tools: Vec<ToolDefinition>,
}
impl SessionContext {
pub fn new(input_messages: Vec<Message>) -> Self {
Self {
input_messages,
..Default::default()
}
}
pub fn add_instructions(&mut self, s: impl Into<String>) {
let s = s.into();
self.instructions = match self.instructions.take() {
Some(existing) => Some(format!("{existing}\n{s}")),
None => Some(s),
};
}
}
#[async_trait]
pub trait ContextProvider: Send + Sync {
async fn before_run(&self, ctx: &mut SessionContext) -> Result<()>;
async fn after_run(
&self,
_request_messages: &[Message],
_response_messages: &[Message],
_error: Option<&Error>,
) -> Result<()> {
Ok(())
}
fn is_history_provider(&self) -> bool {
false
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn add_instructions_sets_when_none() {
let mut ctx = SessionContext::new(vec![]);
assert!(ctx.instructions.is_none());
ctx.add_instructions("be brief");
assert_eq!(ctx.instructions.as_deref(), Some("be brief"));
}
#[test]
fn add_instructions_newline_concatenates() {
let mut ctx = SessionContext::new(vec![]);
ctx.add_instructions("first");
ctx.add_instructions("second");
ctx.add_instructions("third");
assert_eq!(ctx.instructions.as_deref(), Some("first\nsecond\nthird"));
}
#[test]
fn new_sets_input_messages_and_defaults_rest() {
let messages = vec![Message::user("hi")];
let ctx = SessionContext::new(messages.clone());
assert_eq!(ctx.input_messages.len(), messages.len());
assert_eq!(ctx.input_messages[0].text(), "hi");
assert!(ctx.session_id.is_none());
assert!(ctx.service_session_id.is_none());
assert!(ctx.instructions.is_none());
assert!(ctx.messages.is_empty());
assert!(ctx.tools.is_empty());
}
}