use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use futures::StreamExt;
use tracing::Instrument;
use crate::error::ProviderError;
use crate::observe::{AdapterContext, AdapterEnding, AdapterSlot};
use crate::streaming::{Item, Streamed};
use crate::wasm_compat::{WasmBoxedFuture, WasmBoxedStream, WasmCompatSend, WasmCompatSync};
use crate::wire::document::Reassemble;
use crate::wire::{
Call, Capabilities, Decoder, Flow, Mode, Operation, Out, Request, Response, Shared, Wire,
WireEvent,
};
mod dyn_model;
pub(crate) mod http_transport;
mod local;
pub use dyn_model::DynModel;
pub use local::{Local, Step};
#[derive(Clone, Debug, Default, PartialEq)]
pub struct Model<W, T = crate::http_client::DynHttpClient> {
pub wire: W,
pub transport: T,
}
impl<W, T> Model<W, T> {
pub fn new(wire: W, transport: T) -> Self {
Self { wire, transport }
}
}
pub trait Transport<W: Wire>: Clone + WasmCompatSend + WasmCompatSync + 'static {
fn send(&self, payload: W::Payload, exchange: Exchange) -> Opening<W::Frame>;
}
pub struct Exchange {
pub mode: Mode,
pub(crate) observation: Option<AdapterContext>,
}
pub struct Opening<F>(WasmBoxedFuture<'static, Result<Opened<F>, ProviderError>>);
impl<F: WasmCompatSend + 'static> Opening<F> {
pub fn new(
open: impl Future<Output = Result<Opened<F>, ProviderError>> + WasmCompatSend + 'static,
) -> Self {
Self(Box::pin(open))
}
pub fn ready(opened: Opened<F>) -> Self {
Self::new(std::future::ready(Ok(opened)))
}
pub fn failed(error: ProviderError) -> Self {
Self::new(std::future::ready(Err(error)))
}
}
impl<F> Future for Opening<F> {
type Output = Result<Opened<F>, ProviderError>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.0.as_mut().poll(cx)
}
}
pub struct Opened<F> {
pub(crate) frames: WasmBoxedStream<'static, Result<F, ProviderError>>,
pub(crate) request_id: Option<String>,
pub(crate) status: Option<http::StatusCode>,
pub(crate) headers: Option<http::HeaderMap>,
pub(crate) route: Option<String>,
pub(crate) document: Option<serde_json::Value>,
pub(crate) slot: Option<AdapterSlot>,
pub(crate) analysis_only: Option<fn(&F) -> bool>,
}
impl<F: WasmCompatSend + 'static> Opened<F> {
pub fn new(
frames: impl futures::Stream<Item = Result<F, ProviderError>> + WasmCompatSend + 'static,
) -> Self {
Self {
frames: Box::pin(frames),
request_id: None,
status: None,
headers: None,
route: None,
document: None,
slot: None,
analysis_only: None,
}
}
pub fn failed(error: ProviderError) -> Self {
Self::new(futures::stream::once(async move { Err(error) }))
}
pub fn with_request_id(mut self, request_id: Option<String>) -> Self {
self.request_id = crate::provider_response::reported(request_id);
self
}
pub fn with_http(mut self, status: http::StatusCode, headers: http::HeaderMap) -> Self {
self.status = Some(status);
self.headers = Some(headers);
self
}
pub fn with_route(mut self, route: impl Into<String>) -> Self {
self.route = Some(route.into());
self
}
pub fn with_document(mut self, document: serde_json::Value) -> Self {
self.document = Some(document);
self
}
pub fn map_frames<S>(
mut self,
frames: impl FnOnce(WasmBoxedStream<'static, Result<F, ProviderError>>) -> S,
) -> Self
where
S: futures::Stream<Item = Result<F, ProviderError>> + WasmCompatSend + 'static,
{
self.frames = Box::pin(frames(self.frames));
self
}
}
impl<W, T> Model<W, T>
where
W: Wire,
T: Transport<W>,
{
pub fn name(&self) -> &str {
self.wire.describe().name
}
pub fn id(&self) -> Option<&str> {
self.wire.describe().model
}
pub fn capabilities(&self) -> Capabilities {
self.wire.describe().capabilities
}
pub fn call(
&self,
request: impl Into<Request<W>>,
) -> impl Future<Output = Result<Response<W>, ProviderError>> + WasmCompatSend + 'static {
self.finished(request.into(), None)
}
pub fn call_observed(
&self,
request: impl Into<Request<W>>,
observation: AdapterContext,
) -> impl Future<Output = Result<Response<W>, ProviderError>> + WasmCompatSend + 'static {
self.finished(request.into(), Some(observation))
}
fn finished(
&self,
request: Request<W>,
observation: Option<AdapterContext>,
) -> impl Future<Output = Result<Response<W>, ProviderError>> + WasmCompatSend + 'static {
let model = self.clone();
async move {
model
.open(request, Mode::Unary, observation)?
.finish()
.await
}
}
pub fn stream(&self, request: impl Into<Request<W>>) -> Result<Streamed<W::Op>, ProviderError> {
self.open(request.into(), Mode::Streaming, None)
}
pub fn stream_observed(
&self,
request: impl Into<Request<W>>,
observation: AdapterContext,
) -> Result<Streamed<W::Op>, ProviderError> {
self.open(request.into(), Mode::Streaming, Some(observation))
}
pub(crate) async fn call_routed(
&self,
request: Request<W>,
) -> Result<Response<W>, (ProviderError, String)> {
self.open(request, Mode::Unary, None)
.map_err(|error| (error, String::new()))?
.finish_routed()
.await
}
pub(crate) fn open(
&self,
request: Request<W>,
mode: Mode,
observation: Option<AdapterContext>,
) -> Result<Streamed<W::Op>, ProviderError> {
let describe = self.wire.describe();
let request = <W::Op as Operation>::prepare(request, &describe)?;
let provider = describe.name.to_owned();
let mut call = Call::new(&describe, mode);
let fold = <W::Op as Operation>::fold(&request, &mut call);
let span = call.span;
let payload = self.wire.encode(request, mode)?;
let opening = self.transport.send(payload, Exchange { mode, observation });
let shared = Arc::new(Mutex::new(Shared::new(fold)));
let reading = read(
self.wire.clone(),
opening,
Arc::clone(&shared),
span.clone(),
mode,
);
Ok(Streamed::new(reading, shared, span, provider))
}
}
fn read<W: Wire>(
wire: W,
opening: Opening<W::Frame>,
shared: Arc<Mutex<Shared<W::Op>>>,
span: tracing::Span,
mode: Mode,
) -> WasmBoxedStream<'static, ()> {
let decoding = span.clone();
let reading = async_stream::stream! {
let reply: &Mutex<Shared<W::Op>> = &shared;
let opened = match mode {
Mode::Unary => opening.instrument(span.clone()).await,
Mode::Streaming => opening.await,
};
let Opened {
mut frames,
request_id,
status,
headers,
route,
document,
slot,
analysis_only,
} = match opened {
Ok(opened) => opened,
Err(error) => {
fail(reply, None, error);
return;
}
};
record_request_id(&span, request_id.as_deref());
let mut reassembler = document.is_none().then(|| wire.reassembler());
{
let mut state = lock(reply);
state.request_id.clone_from(&request_id);
state.document = document;
state.route = route.unwrap_or_default();
}
let enrich = |error: ProviderError| match mode {
Mode::Unary => error
.with_provider_status(status)
.with_provider_request_id(request_id.clone())
.with_response_headers(headers.clone()),
Mode::Streaming => error,
};
let mut decoder = wire.decoder();
let mut tally = slot.as_ref().map(|slot| Tally {
slot,
analysis_only,
counted: 0,
});
loop {
let flow = match frames.next().await {
Some(Ok(frame)) => step(
&mut decoder,
reassembler.as_mut(),
reply,
frame,
tally.as_mut(),
),
Some(Err(error)) => {
record(reply, reassembler.map(|document| document.finish()));
fail(reply, slot.as_ref(), error);
return;
}
None => eof(&mut decoder, reply, tally.as_ref()),
};
match flow {
Ok(Flow::More) => yield (),
Ok(Flow::Ended(_)) => break,
Err(error) => {
record(reply, reassembler.map(|document| document.finish()));
fail(reply, slot.as_ref(), enrich(error));
return;
}
}
}
record(reply, reassembler.map(|document| document.finish()));
if let Some(slot) = &slot {
slot.finish(AdapterEnding::Terminal);
}
{
let state = lock(reply);
if let Some(document) = state.raw.as_ref().or(state.document.as_ref()) {
crate::providers::internal::trace_json(
crate::providers::internal::LogTarget::Completions,
"reply",
document,
);
}
}
yield ();
};
let mut reading: WasmBoxedStream<'static, ()> = Box::pin(reading);
match mode {
Mode::Streaming => Box::pin(futures::stream::poll_fn(move |cx| {
let _decoding = decoding.enter();
reading.as_mut().poll_next(cx)
})),
Mode::Unary => reading,
}
}
#[cfg(any(test, feature = "websocket", feature = "test-utils"))]
pub(crate) struct Decoded<Op: Operation> {
#[cfg(any(test, feature = "test-utils"))]
pub(crate) items: Vec<Result<crate::streaming::Item<Op::Event>, ProviderError>>,
pub(crate) outcome: Result<Op::Response, ProviderError>,
}
pub(crate) struct Tally<'a, F> {
slot: &'a AdapterSlot,
analysis_only: Option<fn(&F) -> bool>,
counted: usize,
}
pub(crate) fn triage<E>(event: WireEvent<E>) -> Result<Item<E>, ProviderError> {
match event {
WireEvent::Known(event) => Ok(Item::Event(event)),
WireEvent::Unknown { event_type, value } => {
warn_unmodeled(&event_type, &value);
Ok(Item::Unknown(value))
}
WireEvent::Corrupt(error) => Err(ProviderError::from(error)),
}
}
pub(crate) fn step<'id, Op, F, D, R>(
decoder: &mut D,
reassembler: Option<&mut R>,
reply: &'id Mutex<Shared<Op>>,
frame: F,
tally: Option<&mut Tally<'_, F>>,
) -> Result<Flow, ProviderError>
where
Op: Operation,
D: Decoder<'id, Op, F>,
R: Reassemble<F>,
{
if let Some(reassembler) = reassembler {
reassembler.absorb(&frame);
}
let exempt = tally
.as_ref()
.and_then(|tally| tally.analysis_only)
.is_some_and(|analysis_only| analysis_only(&frame));
let classified = decoder.classify(frame);
if let Some(tally) = tally {
let corrupt = matches!(classified, WireEvent::Corrupt(_));
if corrupt || !exempt {
tally.counted += 1;
}
if corrupt {
tally.slot.corrupt(tally.counted);
}
}
match triage(classified)? {
Item::Event(event) => decoder.decode(event, Out::new(reply)),
Item::Unknown(value) => {
Out::new(reply).unknown(value);
Ok(Flow::More)
}
}
}
fn eof<'id, Op, F, D>(
decoder: &mut D,
reply: &'id Mutex<Shared<Op>>,
tally: Option<&Tally<'_, F>>,
) -> Result<Flow, ProviderError>
where
Op: Operation,
D: Decoder<'id, Op, F>,
{
if let Some(tally) = tally {
tally.slot.transport_eof(tally.counted);
}
let step = decoder.eof(Out::new(reply));
if !matches!(step, Ok(Flow::Ended(_)))
&& let Some(tally) = tally
{
tally.slot.eof(tally.counted);
}
match step {
Ok(Flow::More) => Err(ProviderError::Truncated),
step => step,
}
}
#[cfg(any(test, feature = "websocket", feature = "test-utils"))]
pub(crate) fn feed<'id, Op, F, D, R>(
decoder: &mut D,
mut reassembler: Option<R>,
reply: &'id Mutex<Shared<Op>>,
frames: impl IntoIterator<Item = F>,
) -> Result<(), ProviderError>
where
Op: Operation,
D: Decoder<'id, Op, F>,
R: Reassemble<F>,
{
let fed = feed_until_end(decoder, reassembler.as_mut(), reply, frames);
record(reply, reassembler.map(|document| document.finish()));
fed
}
#[cfg(any(test, feature = "websocket", feature = "test-utils"))]
fn feed_until_end<'id, Op, F, D, R>(
decoder: &mut D,
mut reassembler: Option<&mut R>,
reply: &'id Mutex<Shared<Op>>,
frames: impl IntoIterator<Item = F>,
) -> Result<(), ProviderError>
where
Op: Operation,
D: Decoder<'id, Op, F>,
R: Reassemble<F>,
{
for frame in frames {
if let Flow::Ended(_) = step(decoder, reassembler.as_deref_mut(), reply, frame, None)? {
return Ok(());
}
}
eof(decoder, reply, None).map(drop)
}
pub(crate) fn record<Op: Operation>(
reply: &Mutex<Shared<Op>>,
document: Option<serde_json::Value>,
) {
if let Some(document) = document.filter(|document| !document.is_null()) {
lock(reply).raw = Some(document);
}
}
#[cfg(any(test, feature = "websocket", feature = "test-utils"))]
pub(crate) fn settle<Op: Operation>(
shared: Mutex<Shared<Op>>,
fed: Result<(), ProviderError>,
reply: crate::wire::Reply,
) -> Decoded<Op> {
let mut shared = shared
.into_inner()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Err(error) = fed {
shared.items.push_back(Err(error));
}
shared.document = Some(reply.raw).filter(|raw| !raw.is_null());
shared.request_id = reply.provider_request_id;
let mut items = Vec::new();
while let Some(item) = shared.take() {
let failed = item.is_err();
items.push(item);
if failed {
break;
}
}
let outcome = match items.last() {
Some(Err(error)) => Err(error.clone()),
_ => shared.conclude(&reply.provider),
};
#[cfg(any(test, feature = "test-utils"))]
if let Err(error) = &outcome
&& !matches!(items.last(), Some(Err(_)))
{
items.push(Err(error.clone()));
}
Decoded {
#[cfg(any(test, feature = "test-utils"))]
items,
outcome,
}
}
#[cfg(any(test, feature = "test-utils"))]
pub(crate) fn relay_frames<W>(
wire: &W,
frames: impl IntoIterator<Item = W::Frame>,
) -> crate::streaming::StreamEvents
where
W: Wire<Op = crate::operation::Completion>,
{
use crate::error::ErrorReport;
use crate::streaming::Relayed;
let provider = wire.describe().name.to_owned();
let shared = Mutex::new(Shared::new(crate::operation::Turn::relayed(
provider.clone(),
)));
let fed = feed(
&mut wire.decoder(),
Some(wire.reassembler()),
&shared,
frames,
);
let decoded = settle(
shared,
fed,
crate::wire::Reply {
provider,
raw: serde_json::Value::Null,
provider_request_id: None,
},
);
let origin = decoded
.outcome
.as_ref()
.map(|response| response.origin.clone())
.ok();
let relayed: Vec<Result<Relayed, ErrorReport>> = origin
.map(|origin| Ok(Relayed::Origin(origin)))
.into_iter()
.chain(decoded.items.into_iter().map(|item| match item {
Ok(item) => Ok(Relayed::Item(item)),
Err(error) => Err(ErrorReport::from(&error)),
}))
.chain(
decoded
.outcome
.ok()
.map(|response| Ok(Relayed::Done(Box::new(response)))),
)
.collect();
Box::pin(futures::stream::iter(relayed))
}
#[cfg(test)]
impl Decoded<crate::operation::Completion> {
pub(crate) fn events(&self) -> Vec<&crate::streaming::StreamEvent> {
self.items
.iter()
.filter_map(|item| match item {
Ok(crate::streaming::Item::Event(event)) => Some(event),
_ => None,
})
.collect()
}
pub(crate) fn ended(&self) -> Vec<crate::message::AssistantContent> {
self.events()
.into_iter()
.filter_map(|event| match event {
crate::streaming::StreamEvent::End { content, .. } => Some(content.clone()),
_ => None,
})
.collect()
}
}
#[cfg(test)]
macro_rules! decode_events {
($decoder:expr, $provider:expr, $events:expr) => {
$crate::driver::decode_with(
$crate::operation::Turn::relayed($provider),
$provider,
|reply| {
let mut decoder = $decoder;
for event in $events {
if let $crate::wire::Flow::Ended(_) =
$crate::wire::Decoder::decode(&mut decoder, event, reply.out())?
{
return Ok(());
}
}
match $crate::wire::Decoder::eof(&mut decoder, reply.out())? {
$crate::wire::Flow::Ended(_) => Ok(()),
$crate::wire::Flow::More => Err($crate::error::ProviderError::Truncated),
}
},
)
};
}
#[cfg(test)]
pub(crate) use decode_events;
#[cfg(test)]
macro_rules! feed_frames {
($decoder:expr, $provider:expr, $frames:expr) => {
$crate::driver::decode_with(
$crate::operation::Turn::relayed($provider),
$provider,
|reply| {
let mut decoder = $decoder;
reply.feed(
&mut decoder,
None::<$crate::wire::document::Unreassembled>,
$frames,
)
},
)
};
($decoder:expr, $reassembler:expr, $provider:expr, $frames:expr) => {
$crate::driver::decode_with(
$crate::operation::Turn::relayed($provider),
$provider,
|reply| {
let mut decoder = $decoder;
reply.feed(&mut decoder, Some($reassembler), $frames)
},
)
};
}
#[cfg(test)]
pub(crate) use feed_frames;
#[cfg(test)]
pub(crate) struct Replying<'id, Op: Operation>(&'id Mutex<Shared<Op>>);
#[cfg(test)]
impl<'id, Op: Operation> Replying<'id, Op> {
pub(crate) fn out(&self) -> Out<'id, Op> {
Out::new(self.0)
}
pub(crate) fn feed<F, D: Decoder<'id, Op, F>, R: Reassemble<F>>(
&self,
decoder: &mut D,
reassembler: Option<R>,
frames: impl IntoIterator<Item = F>,
) -> Result<(), ProviderError> {
feed(decoder, reassembler, self.0, frames)
}
}
#[cfg(test)]
pub(crate) fn decode_with<Op: Operation>(
fold: Op::Fold,
provider: &str,
run: impl for<'id> FnOnce(Replying<'id, Op>) -> Result<(), ProviderError>,
) -> Decoded<Op> {
let shared = Mutex::new(Shared::new(fold));
let fed = run(Replying(&shared));
settle(
shared,
fed,
crate::wire::Reply {
provider: provider.to_owned(),
raw: serde_json::Value::Null,
provider_request_id: None,
},
)
}
fn fail<Op: Operation>(
reply: &Mutex<Shared<Op>>,
slot: Option<&AdapterSlot>,
error: ProviderError,
) {
if let Some(slot) = slot {
slot.fail(&error);
}
lock(reply).items.push_back(Err(error));
}
pub(crate) fn lock<T>(mutex: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
mutex
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
pub(crate) async fn follow_cursors<P, Fut>(
provider: &str,
operation: &str,
mut page: impl FnMut(Option<String>) -> Fut,
) -> Result<Vec<P>, ProviderError>
where
Fut: Future<Output = Result<(P, Option<String>), ProviderError>>,
{
let mut pages = Vec::new();
let mut cursor = None;
loop {
let (read, next) = page(cursor.clone()).await?;
pages.push(read);
let Some(next) = next else { break };
if cursor.as_deref() == Some(next.as_str()) {
tracing::warn!(
provider,
operation,
pages = pages.len(),
"listing repeated its pagination cursor; returning the pages fetched so far"
);
break;
}
if pages.len() >= MAX_CONTINUATION_PAGES {
tracing::warn!(
provider,
operation,
pages = pages.len(),
"listing hit its page ceiling with a cursor still advancing; returning the pages \
fetched so far"
);
break;
}
cursor = Some(next);
}
Ok(pages)
}
pub fn warn_unmodeled(kind: &str, payload: &impl serde::Serialize) {
tracing::warn!(
kind,
payload_bytes = unknown_payload_bytes(payload),
"skipping unmodeled wire payload"
);
}
fn unknown_payload_bytes(value: &impl serde::Serialize) -> u64 {
struct CountingWriter(u64);
impl std::io::Write for CountingWriter {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.0 += buf.len() as u64;
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
let mut counter = CountingWriter(0);
let _ = serde_json::to_writer(&mut counter, value);
counter.0
}
const MAX_CONTINUATION_PAGES: usize = 1000;
pub(crate) fn record_request_id(span: &tracing::Span, request_id: Option<&str>) {
if let Some(request_id) = request_id
&& !span.is_disabled()
{
span.record(crate::telemetry::PROVIDER_REQUEST_ID_FIELD, request_id);
}
}
#[cfg(test)]
pub(crate) mod tests;