use std::ops::Deref;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use tokio::sync::{mpsc::UnboundedReceiver, oneshot};
use crate::detail::api::{Api, Kvps};
use crate::detail::model::Model;
use crate::detail::session::{run_item_streaming, NativeItemQueue, NativeRequest, NativeSession};
use crate::detail::task::spawn_blocking;
use crate::error::{FoundryLocalError, Result};
use crate::item::Item;
use crate::item_queue::ItemQueue;
use crate::request::{Request, RequestOptions};
use crate::response::Response;
#[derive(Clone)]
pub struct Session {
inner: Arc<NativeSession>,
}
impl Session {
pub async fn new(model: &Model) -> Result<Session> {
let native = model.selected_native().clone();
let inner = spawn_blocking(move || NativeSession::create(&native)).await?;
Ok(Session {
inner: Arc::new(inner),
})
}
pub async fn process_request(&self, request: Request) -> Result<Response> {
let inner = Arc::clone(&self.inner);
let native = Arc::new(NativeRequest::new(Arc::clone(&inner.api))?);
let native_task = Arc::clone(&native);
let handle = tokio::task::spawn_blocking(move || {
let _guard = inner.lock_ops();
populate_native_request(&inner.api, &native_task, &request)?;
let response = inner.process_request(&native_task)?;
Response::from_native(&response)
});
let guard = CancelGuard::new(native);
let joined = handle.await;
guard.disarm();
joined.map_err(|e| FoundryLocalError::Internal {
reason: format!("blocking task join error: {e}"),
})?
}
pub fn process_streaming_request(&self, request: Request) -> ItemStream {
let Request {
items,
input_queue,
options,
} = request;
let option_pairs = options
.as_ref()
.map(RequestOptions::to_pairs)
.unwrap_or_default();
let input_queue = input_queue.map(ItemQueue::into_native);
let (rx, response_rx) =
run_item_streaming(Arc::clone(&self.inner), items, input_queue, option_pairs);
ItemStream {
rx,
response_rx: Some(response_rx),
}
}
pub async fn set_options(&self, options: RequestOptions) -> Result<()> {
let inner = Arc::clone(&self.inner);
spawn_blocking(move || {
let pairs = options.to_pairs();
let kvps = Kvps::from_pairs(Arc::clone(&inner.api), pairs)?;
inner.set_options(kvps.as_ptr())
})
.await
}
pub fn create_input_queue(&self) -> Result<ItemQueue> {
let native = NativeItemQueue::new(Arc::clone(&self.inner.api))?;
Ok(ItemQueue::from_native(Arc::new(native)))
}
}
impl std::fmt::Debug for Session {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Session").finish_non_exhaustive()
}
}
fn populate_native_request(
api: &Arc<Api>,
native: &NativeRequest,
request: &Request,
) -> Result<()> {
for item in &request.items {
native.add_item_value(item)?;
}
if let Some(queue) = &request.input_queue {
native.add_input_queue(queue.native())?;
}
let pairs = request.option_pairs();
if !pairs.is_empty() {
let kvps = Kvps::from_pairs(Arc::clone(api), pairs)?;
native.set_options(kvps.as_ptr())?;
}
Ok(())
}
struct CancelGuard {
native: Arc<NativeRequest>,
armed: bool,
}
impl CancelGuard {
fn new(native: Arc<NativeRequest>) -> Self {
Self {
native,
armed: true,
}
}
fn disarm(mut self) {
self.armed = false;
}
}
impl Drop for CancelGuard {
fn drop(&mut self) {
if self.armed {
self.native.cancel();
}
}
}
fn worker_ended_without_response() -> FoundryLocalError {
FoundryLocalError::Internal {
reason: "streaming response worker ended without a terminal response".into(),
}
}
pub struct ItemStream {
rx: UnboundedReceiver<Result<Item>>,
response_rx: Option<oneshot::Receiver<Result<Response>>>,
}
impl ItemStream {
pub async fn response(&mut self) -> Result<Response> {
if self.response_rx.is_none() {
return Err(FoundryLocalError::Validation {
reason: "the streaming response has already been taken".into(),
});
}
loop {
match self.rx.try_recv() {
Ok(Ok(_)) => continue,
Ok(Err(error)) => return Err(error),
Err(tokio::sync::mpsc::error::TryRecvError::Empty) => {}
Err(tokio::sync::mpsc::error::TryRecvError::Disconnected) => break,
}
let response_rx = self.response_rx.as_mut().expect("checked above");
tokio::select! {
response = &mut *response_rx => {
self.response_rx.take();
while let Ok(item) = self.rx.try_recv() {
item?;
}
return response.map_err(|_| worker_ended_without_response())?;
}
item = self.rx.recv() => {
match item {
Some(Ok(_)) => {}
Some(Err(error)) => return Err(error),
None => break,
}
}
}
}
let response = self
.response_rx
.as_mut()
.expect("checked above")
.await
.map_err(|_| worker_ended_without_response())?;
self.response_rx.take();
response
}
}
impl Unpin for ItemStream {}
impl futures_core::Stream for ItemStream {
type Item = Result<Item>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
self.rx.poll_recv(cx)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ToolDefinition {
pub name: String,
pub description: Option<String>,
pub json_schema: String,
}
impl ToolDefinition {
pub fn new(name: impl Into<String>, json_schema: impl Into<String>) -> Self {
Self {
name: name.into(),
description: None,
json_schema: json_schema.into(),
}
}
pub fn with_description(mut self, description: impl Into<String>) -> Self {
self.description = Some(description.into());
self
}
}
fn validate_session_task(model: &Model, session: &str, allowed: &[&str]) -> Result<()> {
let info = model.info()?;
check_task(info.task.as_deref(), session, allowed)
}
fn check_task(task: Option<&str>, session: &str, allowed: &[&str]) -> Result<()> {
let task = task.unwrap_or("");
if !allowed.contains(&task) {
let expected = allowed
.iter()
.map(|t| format!("'{t}'"))
.collect::<Vec<_>>()
.join(" or ");
return Err(FoundryLocalError::Validation {
reason: format!("{session} requires a model with task {expected}, but got '{task}'."),
});
}
Ok(())
}
#[derive(Clone)]
pub struct ChatSession {
session: Session,
}
impl ChatSession {
pub async fn new(model: &Model) -> Result<ChatSession> {
validate_session_task(
model,
"ChatSession",
&["chat-completion", "vision-language-chat"],
)?;
Ok(ChatSession {
session: Session::new(model).await?,
})
}
pub async fn add_tool_definition(&self, definition: ToolDefinition) -> Result<()> {
let inner = Arc::clone(&self.session.inner);
spawn_blocking(move || {
inner.add_tool_definition(
&definition.name,
definition.description.as_deref(),
&definition.json_schema,
)
})
.await
}
pub async fn remove_tool_definition(&self, name: impl Into<String>) -> Result<bool> {
let inner = Arc::clone(&self.session.inner);
let name = name.into();
spawn_blocking(move || inner.remove_tool_definition(&name)).await
}
pub fn turn_count(&self) -> usize {
self.session.inner.turn_count()
}
pub async fn undo_turns(&self, count: usize) -> Result<()> {
let inner = Arc::clone(&self.session.inner);
spawn_blocking(move || inner.undo_turns(count)).await
}
pub fn into_session(self) -> Session {
self.session
}
}
impl Deref for ChatSession {
type Target = Session;
fn deref(&self) -> &Session {
&self.session
}
}
#[derive(Clone)]
pub struct EmbeddingsSession {
session: Session,
}
impl EmbeddingsSession {
pub async fn new(model: &Model) -> Result<EmbeddingsSession> {
validate_session_task(model, "EmbeddingsSession", &["embeddings"])?;
Ok(EmbeddingsSession {
session: Session::new(model).await?,
})
}
pub async fn embed(&self, input: impl Into<String>) -> Result<Vec<f32>> {
let vectors = self.embed_batch(vec![input.into()]).await?;
vectors
.into_iter()
.next()
.ok_or_else(|| FoundryLocalError::Validation {
reason: "embeddings response contained no vectors".to_string(),
})
}
pub async fn embed_batch(&self, inputs: Vec<String>) -> Result<Vec<Vec<f32>>> {
let items: Vec<Item> = inputs.iter().map(|s| Item::text(s.as_str())).collect();
let request = Request::from_items(items);
let response = self.session.process_request(request).await?;
if response.items.len() != inputs.len() {
return Err(FoundryLocalError::Validation {
reason: format!(
"embeddings response returned {} vectors for {} inputs",
response.items.len(),
inputs.len()
),
});
}
let mut vectors = Vec::with_capacity(response.items.len());
for item in &response.items {
let tensor = item
.as_tensor()
.ok_or_else(|| FoundryLocalError::Validation {
reason: "embeddings response item was not a tensor".to_string(),
})?;
let floats = tensor
.as_f32()
.ok_or_else(|| FoundryLocalError::Validation {
reason: "embeddings tensor was not float data".to_string(),
})?;
vectors.push(floats);
}
Ok(vectors)
}
pub fn into_session(self) -> Session {
self.session
}
}
impl Deref for EmbeddingsSession {
type Target = Session;
fn deref(&self) -> &Session {
&self.session
}
}
#[derive(Clone)]
pub struct AudioSession {
session: Session,
}
impl AudioSession {
pub async fn new(model: &Model) -> Result<AudioSession> {
validate_session_task(model, "AudioSession", &["automatic-speech-recognition"])?;
Ok(AudioSession {
session: Session::new(model).await?,
})
}
pub async fn transcribe(&self, audio: Item) -> Result<String> {
let response = self
.session
.process_request(Request::from_items(vec![audio]))
.await?;
let mut text = String::new();
for item in &response.items {
if let Some(result) = item.as_speech_result() {
text.push_str(&result.text);
} else if let Some(t) = item.as_text() {
text.push_str(t);
}
}
Ok(text)
}
pub fn into_session(self) -> Session {
self.session
}
}
impl Deref for AudioSession {
type Target = Session;
fn deref(&self) -> &Session {
&self.session
}
}
#[cfg(test)]
mod tests {
use std::future::Future;
use std::sync::Arc;
use std::task::{Wake, Waker};
use super::*;
struct NoopWake;
impl Wake for NoopWake {
fn wake(self: Arc<Self>) {}
}
#[tokio::test]
async fn item_stream_returns_terminal_response_once() {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
let (response_tx, response_rx) = oneshot::channel();
let expected = Response {
items: vec![Item::text("complete")],
finish_reason: crate::FinishReason::Stop,
usage: crate::Usage {
prompt_tokens: 3,
completion_tokens: 1,
total_tokens: 4,
},
};
tx.send(Ok(Item::text("chunk one"))).unwrap();
tx.send(Ok(Item::text("chunk two"))).unwrap();
drop(tx);
response_tx.send(Ok(expected.clone())).unwrap();
let mut stream = ItemStream {
rx,
response_rx: Some(response_rx),
};
assert_eq!(stream.response().await.unwrap(), expected);
assert!(stream.rx.is_empty());
assert!(matches!(
stream.response().await,
Err(FoundryLocalError::Validation { .. })
));
}
#[tokio::test]
async fn cancelling_response_await_keeps_terminal_response_available() {
let (_tx, rx) = tokio::sync::mpsc::unbounded_channel();
let (response_tx, response_rx) = oneshot::channel();
let mut stream = ItemStream {
rx,
response_rx: Some(response_rx),
};
let mut response = Box::pin(stream.response());
let waker = Waker::from(Arc::new(NoopWake));
let mut context = Context::from_waker(&waker);
assert!(matches!(
response.as_mut().poll(&mut context),
Poll::Pending
));
drop(response);
response_tx
.send(Ok(Response {
items: Vec::new(),
finish_reason: crate::FinishReason::Stop,
usage: crate::Usage::default(),
}))
.unwrap();
assert_eq!(
stream.response().await.unwrap().finish_reason,
crate::FinishReason::Stop
);
}
#[test]
fn check_task_accepts_allowed_tasks() {
assert!(check_task(
Some("chat-completion"),
"ChatSession",
&["chat-completion", "vision-language-chat"]
)
.is_ok());
assert!(check_task(
Some("vision-language-chat"),
"ChatSession",
&["chat-completion", "vision-language-chat"]
)
.is_ok());
assert!(check_task(Some("embeddings"), "EmbeddingsSession", &["embeddings"]).is_ok());
}
#[test]
fn check_task_rejects_wrong_task() {
let err = check_task(
Some("chat-completion"),
"EmbeddingsSession",
&["embeddings"],
)
.expect_err("wrong task should be rejected");
match err {
FoundryLocalError::Validation { reason } => {
assert!(reason.contains("EmbeddingsSession"), "reason: {reason}");
assert!(reason.contains("'embeddings'"), "reason: {reason}");
assert!(reason.contains("'chat-completion'"), "reason: {reason}");
}
other => panic!("expected Validation error, got: {other:?}"),
}
}
#[test]
fn check_task_treats_missing_task_as_mismatch() {
let err = check_task(None, "AudioSession", &["automatic-speech-recognition"])
.expect_err("missing task should be rejected");
assert!(matches!(err, FoundryLocalError::Validation { .. }));
}
}