use std::rc::Rc;
use perspective_js::utils::{ApiError, ApiResult};
use wasm_bindgen::prelude::*;
use super::client::{AgentTransport, ChatMessage, OnDelta, TurnError, run_turn};
use super::config::{AgentConfig, SystemRole};
use super::docs::DocsCell;
use super::tools::ToolCtx;
use crate::custom_elements::viewer::PerspectiveViewerElement;
const PREAMBLE: &str = include_str!("preamble.md");
const DEFAULT_MAX_TURNS: usize = 16;
pub struct AgentRuntime {
config: AgentConfig,
history: async_lock::Mutex<Vec<ChatMessage>>,
elem: PerspectiveViewerElement,
docs: Option<Rc<DocsCell>>,
}
impl AgentRuntime {
pub fn new(config: &JsValue, elem: PerspectiveViewerElement) -> ApiResult<Self> {
let config = AgentConfig::from_js(config)?;
let docs = config.docs.clone().map(|x| Rc::new(DocsCell::new(x)));
Ok(Self {
config,
history: async_lock::Mutex::new(vec![]),
elem,
docs,
})
}
pub async fn prompt(&self, prompt: String, on_delta: OnDelta<'_>) -> ApiResult<String> {
let mut history = self.history.lock().await;
let transport = self.transport()?;
let bundle = match &self.docs {
Some(docs) => Some(docs.bundle().await),
None => None,
};
let ctx = ToolCtx {
docs: self.docs.clone(),
bundle,
entitlements: self.config.entitlements.clone(),
};
let preamble = self.preamble(&ctx);
let mut messages = Vec::with_capacity(history.len() + 2);
let preamble_len = match self.config.system_role {
SystemRole::System => {
messages.push(ChatMessage::System {
content: preamble.clone(),
});
1
},
SystemRole::User => 0,
};
messages.extend(history.iter().cloned());
let content = if preamble_len == 0 && history.is_empty() {
format!("{preamble}\n\n{prompt}")
} else {
prompt
};
messages.push(ChatMessage::User { content });
let result = run_turn(
&transport,
self.config.model_name(),
self.config.max_turns.unwrap_or(DEFAULT_MAX_TURNS),
&self.elem,
&ctx,
&mut messages,
self.config.system_role,
on_delta,
)
.await;
match result {
Ok(text) => {
*history = messages.split_off(preamble_len);
Ok(text)
},
Err(err @ TurnError::Budget { .. }) => {
*history = messages.split_off(preamble_len);
Err(err.into())
},
Err(err) => Err(err.into()),
}
}
pub async fn reset(&self) {
self.history.lock().await.clear();
}
pub fn label(&self) -> String {
format!(
"{} · {}",
self.config.label_name(),
self.config.model_name()
)
}
fn transport(&self) -> ApiResult<AgentTransport> {
if let Some(engine) = &self.config.engine {
return Ok(AgentTransport::Engine(engine.clone()));
}
let url = self
.config
.url
.clone()
.ok_or_else(|| ApiError::from("`url` or `engine` is required"))?;
let mut headers = self
.config
.headers
.clone()
.map(|x| x.into_iter().collect::<Vec<_>>())
.unwrap_or_default();
if let Some(key) = &self.config.api_key {
headers.push(("authorization".to_owned(), format!("Bearer {key}")));
}
Ok(AgentTransport::Fetch { url, headers })
}
fn preamble(&self, ctx: &ToolCtx) -> String {
let mut preamble = ctx
.bundle
.as_ref()
.and_then(|x| x.as_ref().as_ref().ok())
.and_then(|x| x.preamble.as_deref())
.unwrap_or(PREAMBLE)
.trim_end()
.to_owned();
if let Some(extra) = &self.config.system_prompt {
preamble.push_str("\n\n");
preamble.push_str(extra);
}
preamble
}
}