1use std::future::Future;
22use std::pin::Pin;
23use std::sync::{Arc, Mutex};
24use std::task::{Context, Poll};
25
26use futures::StreamExt;
27use tracing::Instrument;
28
29use crate::error::ProviderError;
30use crate::observe::{AdapterContext, AdapterEnding, AdapterSlot};
31use crate::streaming::Streamed;
32use crate::wasm_compat::{WasmBoxedFuture, WasmBoxedStream, WasmCompatSend, WasmCompatSync};
33use crate::wire::{
34 Call, Capabilities, Decoder, Flow, Mode, Operation, Out, Request, Response, Shared, Wire,
35 WireEvent,
36};
37
38mod dyn_model;
39mod http_transport;
40mod local;
41
42pub use dyn_model::DynModel;
43pub use local::{Local, Step};
44
45#[derive(Clone, Debug, Default, PartialEq)]
52pub struct Model<W, T = crate::http_client::DynHttpClient> {
53 pub wire: W,
55 pub transport: T,
57}
58
59impl<W, T> Model<W, T> {
60 pub fn new(wire: W, transport: T) -> Self {
62 Self { wire, transport }
63 }
64}
65
66pub trait Transport<W: Wire>: Clone + WasmCompatSend + WasmCompatSync + 'static {
73 fn send(&self, payload: W::Payload, exchange: Exchange) -> Opening<W::Frame>;
78}
79
80pub struct Exchange {
82 pub mode: Mode,
84 pub(crate) observation: Option<AdapterContext>,
86}
87
88pub struct Opening<F>(WasmBoxedFuture<'static, Result<Opened<F>, ProviderError>>);
90
91impl<F: WasmCompatSend + 'static> Opening<F> {
92 pub fn new(
94 open: impl Future<Output = Result<Opened<F>, ProviderError>> + WasmCompatSend + 'static,
95 ) -> Self {
96 Self(Box::pin(open))
97 }
98
99 pub fn ready(opened: Opened<F>) -> Self {
101 Self::new(std::future::ready(Ok(opened)))
102 }
103
104 pub fn failed(error: ProviderError) -> Self {
106 Self::new(std::future::ready(Err(error)))
107 }
108}
109
110impl<F> Future for Opening<F> {
111 type Output = Result<Opened<F>, ProviderError>;
112
113 fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
114 self.0.as_mut().poll(cx)
115 }
116}
117
118pub struct Opened<F> {
121 pub(crate) frames: WasmBoxedStream<'static, Result<F, ProviderError>>,
122 pub(crate) request_id: Option<String>,
123 pub(crate) status: Option<http::StatusCode>,
124 pub(crate) headers: Option<http::HeaderMap>,
125 pub(crate) route: Option<String>,
126 pub(crate) document: Option<serde_json::Value>,
127 pub(crate) slot: Option<AdapterSlot>,
128 pub(crate) analysis_only: Option<fn(&F) -> bool>,
129}
130
131impl<F: WasmCompatSend + 'static> Opened<F> {
132 pub fn new(
134 frames: impl futures::Stream<Item = Result<F, ProviderError>> + WasmCompatSend + 'static,
135 ) -> Self {
136 Self {
137 frames: Box::pin(frames),
138 request_id: None,
139 status: None,
140 headers: None,
141 route: None,
142 document: None,
143 slot: None,
144 analysis_only: None,
145 }
146 }
147
148 pub fn failed(error: ProviderError) -> Self {
150 Self::new(futures::stream::once(async move { Err(error) }))
151 }
152
153 pub fn with_request_id(mut self, request_id: Option<String>) -> Self {
156 self.request_id = crate::provider_response::reported(request_id);
157 self
158 }
159
160 pub fn with_http(mut self, status: http::StatusCode, headers: http::HeaderMap) -> Self {
162 self.status = Some(status);
163 self.headers = Some(headers);
164 self
165 }
166
167 pub fn with_route(mut self, route: impl Into<String>) -> Self {
169 self.route = Some(route.into());
170 self
171 }
172
173 pub fn with_document(mut self, document: serde_json::Value) -> Self {
176 self.document = Some(document);
177 self
178 }
179
180 pub fn map_frames<S>(
183 mut self,
184 frames: impl FnOnce(WasmBoxedStream<'static, Result<F, ProviderError>>) -> S,
185 ) -> Self
186 where
187 S: futures::Stream<Item = Result<F, ProviderError>> + WasmCompatSend + 'static,
188 {
189 self.frames = Box::pin(frames(self.frames));
190 self
191 }
192}
193
194impl<W, T> Model<W, T>
195where
196 W: Wire,
197 T: Transport<W>,
198{
199 pub fn name(&self) -> &str {
201 self.wire.describe().name
202 }
203
204 pub fn id(&self) -> Option<&str> {
206 self.wire.describe().model
207 }
208
209 pub fn capabilities(&self) -> Capabilities {
212 self.wire.describe().capabilities
213 }
214
215 pub fn call(
220 &self,
221 request: impl Into<Request<W>>,
222 ) -> impl Future<Output = Result<Response<W>, ProviderError>> + WasmCompatSend + 'static {
223 self.finished(request.into(), None)
224 }
225
226 pub fn call_observed(
228 &self,
229 request: impl Into<Request<W>>,
230 observation: AdapterContext,
231 ) -> impl Future<Output = Result<Response<W>, ProviderError>> + WasmCompatSend + 'static {
232 self.finished(request.into(), Some(observation))
233 }
234
235 fn finished(
238 &self,
239 request: Request<W>,
240 observation: Option<AdapterContext>,
241 ) -> impl Future<Output = Result<Response<W>, ProviderError>> + WasmCompatSend + 'static {
242 let model = self.clone();
243 async move {
244 model
245 .open(request, Mode::Unary, observation)?
246 .finish()
247 .await
248 }
249 }
250
251 pub fn stream(&self, request: impl Into<Request<W>>) -> Result<Streamed<W::Op>, ProviderError> {
255 self.open(request.into(), Mode::Streaming, None)
256 }
257
258 pub fn stream_observed(
260 &self,
261 request: impl Into<Request<W>>,
262 observation: AdapterContext,
263 ) -> Result<Streamed<W::Op>, ProviderError> {
264 self.open(request.into(), Mode::Streaming, Some(observation))
265 }
266
267 pub(crate) async fn call_routed(
270 &self,
271 request: Request<W>,
272 ) -> Result<Response<W>, (ProviderError, String)> {
273 self.open(request, Mode::Unary, None)
274 .map_err(|error| (error, String::new()))?
275 .finish_routed()
276 .await
277 }
278
279 pub(crate) fn open(
283 &self,
284 request: Request<W>,
285 mode: Mode,
286 observation: Option<AdapterContext>,
287 ) -> Result<Streamed<W::Op>, ProviderError> {
288 <W::Op as Operation>::validate(&request)?;
289 let describe = self.wire.describe();
290 let provider = describe.name.to_owned();
291 let mut call = Call::new(&describe, mode);
292 let fold = <W::Op as Operation>::fold(&request, &mut call);
293 let span = call.span;
294 let payload = self.wire.encode(request, mode)?;
295 let opening = self.transport.send(payload, Exchange { mode, observation });
296 let shared = Arc::new(Mutex::new(Shared::new(fold)));
297 let reading = read(
298 self.wire.clone(),
299 opening,
300 Arc::clone(&shared),
301 span.clone(),
302 mode,
303 );
304 Ok(Streamed::new(reading, shared, span, provider))
305 }
306}
307
308fn read<W: Wire>(
316 wire: W,
317 opening: Opening<W::Frame>,
318 shared: Arc<Mutex<Shared<W::Op>>>,
319 span: tracing::Span,
320 mode: Mode,
321) -> WasmBoxedStream<'static, ()> {
322 let decoding = span.clone();
323 let reading = async_stream::stream! {
324 let reply: &Mutex<Shared<W::Op>> = &shared;
325 let opened = match mode {
327 Mode::Unary => opening.instrument(span.clone()).await,
328 Mode::Streaming => opening.await,
329 };
330 let Opened {
331 mut frames,
332 request_id,
333 status,
334 headers,
335 route,
336 document,
337 slot,
338 analysis_only,
339 } = match opened {
340 Ok(opened) => opened,
341 Err(error) => {
342 fail(reply, slot_none(), error);
343 return;
344 }
345 };
346 record_request_id(&span, request_id.as_deref());
347 {
348 let mut state = lock(reply);
349 state.request_id.clone_from(&request_id);
350 state.document = document;
351 state.route = route.unwrap_or_default();
352 }
353 let enrich = |error: ProviderError| match mode {
354 Mode::Unary => error
356 .with_provider_status(status)
357 .with_provider_request_id(request_id.clone())
358 .with_response_headers(headers.clone()),
359 Mode::Streaming => error,
360 };
361 let mut decoder = wire.decoder();
362 let mut counted = 0usize;
364 loop {
365 let step = match frames.next().await {
366 Some(Ok(frame)) => {
367 let analysis = slot.is_some()
368 && analysis_only.is_some_and(|analysis_only| analysis_only(&frame));
369 let classified = decoder.classify(frame);
370 let corrupt = matches!(classified, WireEvent::Corrupt(_));
372 if slot.is_some() && (corrupt || !analysis) {
373 counted += 1;
374 }
375 match classified {
376 WireEvent::Known(event) => decoder.decode(event, Out::new(reply)),
377 WireEvent::Unknown { event_type, value } => {
380 warn_unmodeled(&event_type, &value);
381 Out::new(reply).unknown(value);
382 Ok(Flow::More)
383 }
384 WireEvent::Corrupt(error) => {
385 if let Some(slot) = &slot {
386 slot.corrupt(counted);
387 }
388 Err(ProviderError::from(error))
389 }
390 }
391 }
392 Some(Err(error)) => {
393 fail(reply, slot.as_ref(), error);
394 return;
395 }
396 None => {
397 if let Some(slot) = &slot {
398 slot.transport_eof(counted);
399 }
400 let step = decoder.eof(Out::new(reply));
403 if !matches!(step, Ok(Flow::Ended(_)))
404 && let Some(slot) = &slot
405 {
406 slot.eof(counted);
407 }
408 match step {
409 Ok(Flow::More) => Err(ProviderError::Truncated),
410 step => step,
411 }
412 }
413 };
414 match step {
415 Ok(Flow::More) => yield (),
416 Ok(Flow::Ended(_)) => break,
417 Err(error) => {
418 fail(reply, slot.as_ref(), enrich(error));
419 return;
420 }
421 }
422 }
423 if let Some(slot) = &slot {
424 slot.finish(AdapterEnding::Terminal);
425 }
426 {
427 let state = lock(reply);
428 if let Some(document) = state.raw.as_ref().or(state.document.as_ref()) {
429 crate::providers::internal::trace_json(
430 crate::providers::internal::LogTarget::Completions,
431 "reply",
432 document,
433 );
434 }
435 }
436 yield ();
437 };
438 let mut reading: WasmBoxedStream<'static, ()> = Box::pin(reading);
439 match mode {
440 Mode::Streaming => Box::pin(futures::stream::poll_fn(move |cx| {
442 let _decoding = decoding.enter();
443 reading.as_mut().poll_next(cx)
444 })),
445 Mode::Unary => reading,
446 }
447}
448
449#[cfg(any(test, feature = "websocket", feature = "test-utils"))]
452pub(crate) struct Decoded<Op: Operation> {
453 #[cfg(any(test, feature = "test-utils"))]
455 pub(crate) items: Vec<Result<crate::streaming::Item<Op::Event>, ProviderError>>,
456 pub(crate) outcome: Result<Op::Response, ProviderError>,
457}
458
459#[cfg(any(test, feature = "websocket", feature = "test-utils"))]
463pub(crate) fn triage<E>(event: WireEvent<E>) -> Result<crate::streaming::Item<E>, ProviderError> {
464 match event {
465 WireEvent::Known(event) => Ok(crate::streaming::Item::Event(event)),
466 WireEvent::Unknown { event_type, value } => {
467 warn_unmodeled(&event_type, &value);
468 Ok(crate::streaming::Item::Unknown(value))
469 }
470 WireEvent::Corrupt(error) => Err(ProviderError::from(error)),
471 }
472}
473
474#[cfg(any(test, feature = "websocket", feature = "test-utils"))]
476pub(crate) fn step<'id, Op, F, D>(
477 decoder: &mut D,
478 reply: &'id Mutex<Shared<Op>>,
479 frame: F,
480) -> Result<Flow, ProviderError>
481where
482 Op: Operation,
483 D: crate::wire::Decoder<'id, Op, F>,
484{
485 match triage(decoder.classify(frame))? {
486 crate::streaming::Item::Event(event) => decoder.decode(event, Out::new(reply)),
487 crate::streaming::Item::Unknown(value) => {
488 Out::new(reply).unknown(value);
489 Ok(Flow::More)
490 }
491 }
492}
493
494#[cfg(any(test, feature = "websocket", feature = "test-utils"))]
497fn feed<'id, Op, F, D>(
498 decoder: &mut D,
499 reply: &'id Mutex<Shared<Op>>,
500 frames: impl IntoIterator<Item = F>,
501) -> Result<(), ProviderError>
502where
503 Op: Operation,
504 D: crate::wire::Decoder<'id, Op, F>,
505{
506 for frame in frames {
507 if let Flow::Ended(_) = step(decoder, reply, frame)? {
508 return Ok(());
509 }
510 }
511 match decoder.eof(Out::new(reply))? {
512 Flow::Ended(_) => Ok(()),
513 Flow::More => Err(ProviderError::Truncated),
514 }
515}
516
517#[cfg(any(test, feature = "websocket", feature = "test-utils"))]
520pub(crate) fn settle<Op: Operation>(
521 shared: Mutex<Shared<Op>>,
522 fed: Result<(), ProviderError>,
523 reply: crate::wire::Reply,
524) -> Decoded<Op> {
525 let Shared {
526 mut fold,
527 items,
528 end,
529 raw,
530 ..
531 } = shared
532 .into_inner()
533 .unwrap_or_else(std::sync::PoisonError::into_inner);
534 let mut absorbed = Ok(());
535 for item in &items {
536 if let (Ok(crate::streaming::Item::Event(event)), Ok(())) = (item, &absorbed) {
537 absorbed = crate::wire::Fold::absorb(&mut fold, event);
538 }
539 }
540 let outcome = fed.and(absorbed).and_then(|()| {
541 let reply = crate::wire::Reply {
542 raw: if reply.raw.is_null() {
543 raw.unwrap_or(serde_json::Value::Null)
544 } else {
545 reply.raw
546 },
547 ..reply
548 };
549 crate::wire::Fold::finish(fold, end.ok_or(ProviderError::Truncated)?, reply)
550 });
551 #[cfg(any(test, feature = "test-utils"))]
552 let items = {
553 let mut items: Vec<_> = items.into_iter().collect();
554 if let Err(error) = &outcome
555 && !items.iter().any(Result::is_err)
556 {
557 items.push(Err(error.clone()));
558 }
559 items
560 };
561 Decoded {
562 #[cfg(any(test, feature = "test-utils"))]
563 items,
564 outcome,
565 }
566}
567
568#[cfg(any(test, feature = "websocket"))]
573pub(crate) fn decode_frames<W: Wire>(
574 wire: &W,
575 fold: <W::Op as Operation>::Fold,
576 frames: impl IntoIterator<Item = W::Frame>,
577 reply: crate::wire::Reply,
578) -> Result<Response<W>, ProviderError> {
579 let shared = Mutex::new(Shared::new(fold));
580 let fed = feed(&mut wire.decoder(), &shared, frames);
581 settle(shared, fed, reply).outcome
582}
583
584#[cfg(any(test, feature = "websocket"))]
587pub(crate) fn decode_body<W: Wire<Frame = crate::wire::WireFrame>>(
588 wire: &W,
589 fold: <W::Op as Operation>::Fold,
590 body: String,
591 reply: crate::wire::Reply,
592) -> Result<Response<W>, ProviderError> {
593 decode_frames(wire, fold, [crate::wire::WireFrame::Text(body)], reply)
594}
595
596#[cfg(any(test, feature = "test-utils"))]
599pub(crate) fn relay_frames<W>(
600 wire: &W,
601 frames: impl IntoIterator<Item = W::Frame>,
602) -> crate::streaming::StreamEvents
603where
604 W: Wire<Op = crate::operation::Completion>,
605{
606 use crate::error::ErrorReport;
607 use crate::streaming::Relayed;
608
609 let provider = wire.describe().name.to_owned();
610 let shared = Mutex::new(Shared::new(crate::operation::Turn::new(provider.clone())));
611 let fed = feed(&mut wire.decoder(), &shared, frames);
612 let decoded = settle(
613 shared,
614 fed,
615 crate::wire::Reply {
616 provider,
617 raw: serde_json::Value::Null,
618 provider_request_id: None,
619 },
620 );
621 let relayed: Vec<Result<Relayed, ErrorReport>> = decoded
622 .items
623 .into_iter()
624 .map(|item| match item {
625 Ok(item) => Ok(Relayed::Item(item)),
626 Err(error) => Err(ErrorReport::from(&error)),
627 })
628 .chain(
630 decoded
631 .outcome
632 .ok()
633 .map(|response| Ok(Relayed::Done(Box::new(response)))),
634 )
635 .collect();
636 Box::pin(futures::stream::iter(relayed))
637}
638
639#[cfg(test)]
640impl Decoded<crate::operation::Completion> {
641 pub(crate) fn events(&self) -> Vec<&crate::streaming::StreamEvent> {
643 self.items
644 .iter()
645 .filter_map(|item| match item {
646 Ok(crate::streaming::Item::Event(event)) => Some(event),
647 _ => None,
648 })
649 .collect()
650 }
651
652 pub(crate) fn ended(&self) -> Vec<crate::message::AssistantContent> {
654 self.events()
655 .into_iter()
656 .filter_map(|event| match event {
657 crate::streaming::StreamEvent::End { content, .. } => Some(content.clone()),
658 _ => None,
659 })
660 .collect()
661 }
662}
663
664#[cfg(test)]
667macro_rules! decode_events {
668 ($decoder:expr, $provider:expr, $events:expr) => {
669 $crate::driver::decode_with(
670 $crate::operation::Turn::new($provider),
671 $provider,
672 |reply| {
673 let mut decoder = $decoder;
674 for event in $events {
675 if let $crate::wire::Flow::Ended(_) =
676 $crate::wire::Decoder::decode(&mut decoder, event, reply.out())?
677 {
678 return Ok(());
679 }
680 }
681 match $crate::wire::Decoder::eof(&mut decoder, reply.out())? {
682 $crate::wire::Flow::Ended(_) => Ok(()),
683 $crate::wire::Flow::More => Err($crate::error::ProviderError::Truncated),
684 }
685 },
686 )
687 };
688}
689#[cfg(test)]
690pub(crate) use decode_events;
691
692#[cfg(test)]
696macro_rules! feed_frames {
697 ($decoder:expr, $provider:expr, $frames:expr) => {
698 $crate::driver::decode_with(
699 $crate::operation::Turn::new($provider),
700 $provider,
701 |reply| {
702 let mut decoder = $decoder;
703 reply.feed(&mut decoder, $frames)
704 },
705 )
706 };
707}
708#[cfg(test)]
709pub(crate) use feed_frames;
710
711#[cfg(test)]
713pub(crate) struct Replying<'id, Op: Operation>(&'id Mutex<Shared<Op>>);
714
715#[cfg(test)]
716impl<'id, Op: Operation> Replying<'id, Op> {
717 pub(crate) fn out(&self) -> Out<'id, Op> {
719 Out::new(self.0)
720 }
721
722 pub(crate) fn feed<F, D: crate::wire::Decoder<'id, Op, F>>(
724 &self,
725 decoder: &mut D,
726 frames: impl IntoIterator<Item = F>,
727 ) -> Result<(), ProviderError> {
728 feed(decoder, self.0, frames)
729 }
730}
731
732#[cfg(test)]
736pub(crate) fn decode_with<Op: Operation>(
737 fold: Op::Fold,
738 provider: &str,
739 run: impl for<'id> FnOnce(Replying<'id, Op>) -> Result<(), ProviderError>,
740) -> Decoded<Op> {
741 let shared = Mutex::new(Shared::new(fold));
742 let fed = run(Replying(&shared));
743 settle(
744 shared,
745 fed,
746 crate::wire::Reply {
747 provider: provider.to_owned(),
748 raw: serde_json::Value::Null,
749 provider_request_id: None,
750 },
751 )
752}
753
754fn slot_none() -> Option<&'static AdapterSlot> {
756 None
757}
758
759fn fail<Op: Operation>(
761 reply: &Mutex<Shared<Op>>,
762 slot: Option<&AdapterSlot>,
763 error: ProviderError,
764) {
765 if let Some(slot) = slot {
766 slot.fail(&error);
767 }
768 lock(reply).items.push_back(Err(error));
769}
770
771pub(crate) fn lock<T>(mutex: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
772 mutex
773 .lock()
774 .unwrap_or_else(std::sync::PoisonError::into_inner)
775}
776
777pub(crate) async fn follow_cursors<P, Fut>(
781 provider: &str,
782 operation: &str,
783 mut page: impl FnMut(Option<String>) -> Fut,
784) -> Result<Vec<P>, ProviderError>
785where
786 Fut: Future<Output = Result<(P, Option<String>), ProviderError>>,
787{
788 let mut pages = Vec::new();
789 let mut cursor = None;
790 loop {
791 let (read, next) = page(cursor.clone()).await?;
792 pages.push(read);
793 let Some(next) = next else { break };
796 if cursor.as_deref() == Some(next.as_str()) {
797 tracing::warn!(
800 provider,
801 operation,
802 pages = pages.len(),
803 "listing repeated its pagination cursor; returning the pages fetched so far"
804 );
805 break;
806 }
807 if pages.len() >= MAX_CONTINUATION_PAGES {
808 tracing::warn!(
809 provider,
810 operation,
811 pages = pages.len(),
812 "listing hit its page ceiling with a cursor still advancing; returning the pages \
813 fetched so far"
814 );
815 break;
816 }
817 cursor = Some(next);
818 }
819 Ok(pages)
820}
821
822pub fn warn_unmodeled(kind: &str, payload: &impl serde::Serialize) {
825 tracing::warn!(
826 kind,
827 payload_bytes = unknown_payload_bytes(payload),
828 "skipping unmodeled wire payload"
829 );
830}
831
832fn unknown_payload_bytes(value: &impl serde::Serialize) -> u64 {
835 struct CountingWriter(u64);
838
839 impl std::io::Write for CountingWriter {
840 fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
841 self.0 += buf.len() as u64;
842 Ok(buf.len())
843 }
844
845 fn flush(&mut self) -> std::io::Result<()> {
846 Ok(())
847 }
848 }
849
850 let mut counter = CountingWriter(0);
851 let _ = serde_json::to_writer(&mut counter, value);
853 counter.0
854}
855
856const MAX_CONTINUATION_PAGES: usize = 1000;
859
860pub(crate) fn record_request_id(span: &tracing::Span, request_id: Option<&str>) {
862 if let Some(request_id) = request_id
863 && !span.is_disabled()
864 {
865 span.record(crate::telemetry::PROVIDER_REQUEST_ID_FIELD, request_id);
866 }
867}
868
869#[cfg(test)]
870pub(crate) mod tests;