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::{Item, Streamed};
32use crate::wasm_compat::{WasmBoxedFuture, WasmBoxedStream, WasmCompatSend, WasmCompatSync};
33use crate::wire::document::Reassemble;
34use crate::wire::{
35 Call, Capabilities, Decoder, Flow, Mode, Operation, Out, Request, Response, Shared, Wire,
36 WireEvent,
37};
38
39mod dyn_model;
40pub(crate) mod http_transport;
41mod local;
42
43pub use dyn_model::DynModel;
44pub use local::{Local, Step};
45
46#[derive(Clone, Debug, Default, PartialEq)]
53pub struct Model<W, T = crate::http_client::DynHttpClient> {
54 pub wire: W,
56 pub transport: T,
58}
59
60impl<W, T> Model<W, T> {
61 pub fn new(wire: W, transport: T) -> Self {
63 Self { wire, transport }
64 }
65}
66
67pub trait Transport<W: Wire>: Clone + WasmCompatSend + WasmCompatSync + 'static {
74 fn send(&self, payload: W::Payload, exchange: Exchange) -> Opening<W::Frame>;
79}
80
81pub struct Exchange {
83 pub mode: Mode,
85 pub(crate) observation: Option<AdapterContext>,
87}
88
89pub struct Opening<F>(WasmBoxedFuture<'static, Result<Opened<F>, ProviderError>>);
91
92impl<F: WasmCompatSend + 'static> Opening<F> {
93 pub fn new(
95 open: impl Future<Output = Result<Opened<F>, ProviderError>> + WasmCompatSend + 'static,
96 ) -> Self {
97 Self(Box::pin(open))
98 }
99
100 pub fn ready(opened: Opened<F>) -> Self {
102 Self::new(std::future::ready(Ok(opened)))
103 }
104
105 pub fn failed(error: ProviderError) -> Self {
107 Self::new(std::future::ready(Err(error)))
108 }
109}
110
111impl<F> Future for Opening<F> {
112 type Output = Result<Opened<F>, ProviderError>;
113
114 fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
115 self.0.as_mut().poll(cx)
116 }
117}
118
119pub struct Opened<F> {
122 pub(crate) frames: WasmBoxedStream<'static, Result<F, ProviderError>>,
123 pub(crate) request_id: Option<String>,
124 pub(crate) status: Option<http::StatusCode>,
125 pub(crate) headers: Option<http::HeaderMap>,
126 pub(crate) route: Option<String>,
127 pub(crate) document: Option<serde_json::Value>,
128 pub(crate) slot: Option<AdapterSlot>,
129 pub(crate) analysis_only: Option<fn(&F) -> bool>,
130}
131
132impl<F: WasmCompatSend + 'static> Opened<F> {
133 pub fn new(
135 frames: impl futures::Stream<Item = Result<F, ProviderError>> + WasmCompatSend + 'static,
136 ) -> Self {
137 Self {
138 frames: Box::pin(frames),
139 request_id: None,
140 status: None,
141 headers: None,
142 route: None,
143 document: None,
144 slot: None,
145 analysis_only: None,
146 }
147 }
148
149 pub fn failed(error: ProviderError) -> Self {
151 Self::new(futures::stream::once(async move { Err(error) }))
152 }
153
154 pub fn with_request_id(mut self, request_id: Option<String>) -> Self {
157 self.request_id = crate::provider_response::reported(request_id);
158 self
159 }
160
161 pub fn with_http(mut self, status: http::StatusCode, headers: http::HeaderMap) -> Self {
163 self.status = Some(status);
164 self.headers = Some(headers);
165 self
166 }
167
168 pub fn with_route(mut self, route: impl Into<String>) -> Self {
170 self.route = Some(route.into());
171 self
172 }
173
174 pub fn with_document(mut self, document: serde_json::Value) -> Self {
177 self.document = Some(document);
178 self
179 }
180
181 pub fn map_frames<S>(
184 mut self,
185 frames: impl FnOnce(WasmBoxedStream<'static, Result<F, ProviderError>>) -> S,
186 ) -> Self
187 where
188 S: futures::Stream<Item = Result<F, ProviderError>> + WasmCompatSend + 'static,
189 {
190 self.frames = Box::pin(frames(self.frames));
191 self
192 }
193}
194
195impl<W, T> Model<W, T>
196where
197 W: Wire,
198 T: Transport<W>,
199{
200 pub fn name(&self) -> &str {
202 self.wire.describe().name
203 }
204
205 pub fn id(&self) -> Option<&str> {
207 self.wire.describe().model
208 }
209
210 pub fn capabilities(&self) -> Capabilities {
213 self.wire.describe().capabilities
214 }
215
216 pub fn call(
221 &self,
222 request: impl Into<Request<W>>,
223 ) -> impl Future<Output = Result<Response<W>, ProviderError>> + WasmCompatSend + 'static {
224 self.finished(request.into(), None)
225 }
226
227 pub fn call_observed(
229 &self,
230 request: impl Into<Request<W>>,
231 observation: AdapterContext,
232 ) -> impl Future<Output = Result<Response<W>, ProviderError>> + WasmCompatSend + 'static {
233 self.finished(request.into(), Some(observation))
234 }
235
236 fn finished(
239 &self,
240 request: Request<W>,
241 observation: Option<AdapterContext>,
242 ) -> impl Future<Output = Result<Response<W>, ProviderError>> + WasmCompatSend + 'static {
243 let model = self.clone();
244 async move {
245 model
246 .open(request, Mode::Unary, observation)?
247 .finish()
248 .await
249 }
250 }
251
252 pub fn stream(&self, request: impl Into<Request<W>>) -> Result<Streamed<W::Op>, ProviderError> {
256 self.open(request.into(), Mode::Streaming, None)
257 }
258
259 pub fn stream_observed(
261 &self,
262 request: impl Into<Request<W>>,
263 observation: AdapterContext,
264 ) -> Result<Streamed<W::Op>, ProviderError> {
265 self.open(request.into(), Mode::Streaming, Some(observation))
266 }
267
268 pub(crate) async fn call_routed(
271 &self,
272 request: Request<W>,
273 ) -> Result<Response<W>, (ProviderError, String)> {
274 self.open(request, Mode::Unary, None)
275 .map_err(|error| (error, String::new()))?
276 .finish_routed()
277 .await
278 }
279
280 pub(crate) fn open(
284 &self,
285 request: Request<W>,
286 mode: Mode,
287 observation: Option<AdapterContext>,
288 ) -> Result<Streamed<W::Op>, ProviderError> {
289 let describe = self.wire.describe();
290 let request = <W::Op as Operation>::prepare(request, &describe)?;
291 let provider = describe.name.to_owned();
292 let mut call = Call::new(&describe, mode);
293 let fold = <W::Op as Operation>::fold(&request, &mut call);
294 let span = call.span;
295 let payload = self.wire.encode(request, mode)?;
296 let opening = self.transport.send(payload, Exchange { mode, observation });
297 let shared = Arc::new(Mutex::new(Shared::new(fold)));
298 let reading = read(
299 self.wire.clone(),
300 opening,
301 Arc::clone(&shared),
302 span.clone(),
303 mode,
304 );
305 Ok(Streamed::new(reading, shared, span, provider))
306 }
307}
308
309fn read<W: Wire>(
319 wire: W,
320 opening: Opening<W::Frame>,
321 shared: Arc<Mutex<Shared<W::Op>>>,
322 span: tracing::Span,
323 mode: Mode,
324) -> WasmBoxedStream<'static, ()> {
325 let decoding = span.clone();
326 let reading = async_stream::stream! {
327 let reply: &Mutex<Shared<W::Op>> = &shared;
328 let opened = match mode {
330 Mode::Unary => opening.instrument(span.clone()).await,
331 Mode::Streaming => opening.await,
332 };
333 let Opened {
334 mut frames,
335 request_id,
336 status,
337 headers,
338 route,
339 document,
340 slot,
341 analysis_only,
342 } = match opened {
343 Ok(opened) => opened,
344 Err(error) => {
345 fail(reply, None, error);
346 return;
347 }
348 };
349 record_request_id(&span, request_id.as_deref());
350 let mut reassembler = document.is_none().then(|| wire.reassembler());
351 {
352 let mut state = lock(reply);
353 state.request_id.clone_from(&request_id);
354 state.document = document;
355 state.route = route.unwrap_or_default();
356 }
357 let enrich = |error: ProviderError| match mode {
358 Mode::Unary => error
360 .with_provider_status(status)
361 .with_provider_request_id(request_id.clone())
362 .with_response_headers(headers.clone()),
363 Mode::Streaming => error,
364 };
365 let mut decoder = wire.decoder();
366 let mut tally = slot.as_ref().map(|slot| Tally {
367 slot,
368 analysis_only,
369 counted: 0,
370 });
371 loop {
372 let flow = match frames.next().await {
373 Some(Ok(frame)) => step(
374 &mut decoder,
375 reassembler.as_mut(),
376 reply,
377 frame,
378 tally.as_mut(),
379 ),
380 Some(Err(error)) => {
381 record(reply, reassembler.map(|document| document.finish()));
382 fail(reply, slot.as_ref(), error);
383 return;
384 }
385 None => eof(&mut decoder, reply, tally.as_ref()),
386 };
387 match flow {
388 Ok(Flow::More) => yield (),
389 Ok(Flow::Ended(_)) => break,
390 Err(error) => {
391 record(reply, reassembler.map(|document| document.finish()));
392 fail(reply, slot.as_ref(), enrich(error));
393 return;
394 }
395 }
396 }
397 record(reply, reassembler.map(|document| document.finish()));
398 if let Some(slot) = &slot {
399 slot.finish(AdapterEnding::Terminal);
400 }
401 {
402 let state = lock(reply);
403 if let Some(document) = state.raw.as_ref().or(state.document.as_ref()) {
404 crate::providers::internal::trace_json(
405 crate::providers::internal::LogTarget::Completions,
406 "reply",
407 document,
408 );
409 }
410 }
411 yield ();
412 };
413 let mut reading: WasmBoxedStream<'static, ()> = Box::pin(reading);
414 match mode {
415 Mode::Streaming => Box::pin(futures::stream::poll_fn(move |cx| {
417 let _decoding = decoding.enter();
418 reading.as_mut().poll_next(cx)
419 })),
420 Mode::Unary => reading,
421 }
422}
423
424#[cfg(any(test, feature = "websocket", feature = "test-utils"))]
427pub(crate) struct Decoded<Op: Operation> {
428 #[cfg(any(test, feature = "test-utils"))]
430 pub(crate) items: Vec<Result<crate::streaming::Item<Op::Event>, ProviderError>>,
431 pub(crate) outcome: Result<Op::Response, ProviderError>,
432}
433
434pub(crate) struct Tally<'a, F> {
437 slot: &'a AdapterSlot,
438 analysis_only: Option<fn(&F) -> bool>,
440 counted: usize,
441}
442
443pub(crate) fn triage<E>(event: WireEvent<E>) -> Result<Item<E>, ProviderError> {
448 match event {
449 WireEvent::Known(event) => Ok(Item::Event(event)),
450 WireEvent::Unknown { event_type, value } => {
451 warn_unmodeled(&event_type, &value);
452 Ok(Item::Unknown(value))
453 }
454 WireEvent::Corrupt(error) => Err(ProviderError::from(error)),
455 }
456}
457
458pub(crate) fn step<'id, Op, F, D, R>(
461 decoder: &mut D,
462 reassembler: Option<&mut R>,
463 reply: &'id Mutex<Shared<Op>>,
464 frame: F,
465 tally: Option<&mut Tally<'_, F>>,
466) -> Result<Flow, ProviderError>
467where
468 Op: Operation,
469 D: Decoder<'id, Op, F>,
470 R: Reassemble<F>,
471{
472 if let Some(reassembler) = reassembler {
473 reassembler.absorb(&frame);
474 }
475 let exempt = tally
476 .as_ref()
477 .and_then(|tally| tally.analysis_only)
478 .is_some_and(|analysis_only| analysis_only(&frame));
479 let classified = decoder.classify(frame);
480 if let Some(tally) = tally {
481 let corrupt = matches!(classified, WireEvent::Corrupt(_));
483 if corrupt || !exempt {
484 tally.counted += 1;
485 }
486 if corrupt {
487 tally.slot.corrupt(tally.counted);
488 }
489 }
490 match triage(classified)? {
491 Item::Event(event) => decoder.decode(event, Out::new(reply)),
492 Item::Unknown(value) => {
493 Out::new(reply).unknown(value);
494 Ok(Flow::More)
495 }
496 }
497}
498
499fn eof<'id, Op, F, D>(
502 decoder: &mut D,
503 reply: &'id Mutex<Shared<Op>>,
504 tally: Option<&Tally<'_, F>>,
505) -> Result<Flow, ProviderError>
506where
507 Op: Operation,
508 D: Decoder<'id, Op, F>,
509{
510 if let Some(tally) = tally {
511 tally.slot.transport_eof(tally.counted);
512 }
513 let step = decoder.eof(Out::new(reply));
514 if !matches!(step, Ok(Flow::Ended(_)))
515 && let Some(tally) = tally
516 {
517 tally.slot.eof(tally.counted);
518 }
519 match step {
520 Ok(Flow::More) => Err(ProviderError::Truncated),
521 step => step,
522 }
523}
524
525#[cfg(any(test, feature = "websocket", feature = "test-utils"))]
529pub(crate) fn feed<'id, Op, F, D, R>(
530 decoder: &mut D,
531 mut reassembler: Option<R>,
532 reply: &'id Mutex<Shared<Op>>,
533 frames: impl IntoIterator<Item = F>,
534) -> Result<(), ProviderError>
535where
536 Op: Operation,
537 D: Decoder<'id, Op, F>,
538 R: Reassemble<F>,
539{
540 let fed = feed_until_end(decoder, reassembler.as_mut(), reply, frames);
541 record(reply, reassembler.map(|document| document.finish()));
542 fed
543}
544
545#[cfg(any(test, feature = "websocket", feature = "test-utils"))]
546fn feed_until_end<'id, Op, F, D, R>(
547 decoder: &mut D,
548 mut reassembler: Option<&mut R>,
549 reply: &'id Mutex<Shared<Op>>,
550 frames: impl IntoIterator<Item = F>,
551) -> Result<(), ProviderError>
552where
553 Op: Operation,
554 D: Decoder<'id, Op, F>,
555 R: Reassemble<F>,
556{
557 for frame in frames {
558 if let Flow::Ended(_) = step(decoder, reassembler.as_deref_mut(), reply, frame, None)? {
559 return Ok(());
560 }
561 }
562 eof(decoder, reply, None).map(drop)
563}
564
565pub(crate) fn record<Op: Operation>(
568 reply: &Mutex<Shared<Op>>,
569 document: Option<serde_json::Value>,
570) {
571 if let Some(document) = document.filter(|document| !document.is_null()) {
572 lock(reply).raw = Some(document);
573 }
574}
575
576#[cfg(any(test, feature = "websocket", feature = "test-utils"))]
580pub(crate) fn settle<Op: Operation>(
581 shared: Mutex<Shared<Op>>,
582 fed: Result<(), ProviderError>,
583 reply: crate::wire::Reply,
584) -> Decoded<Op> {
585 let mut shared = shared
586 .into_inner()
587 .unwrap_or_else(std::sync::PoisonError::into_inner);
588 if let Err(error) = fed {
589 shared.items.push_back(Err(error));
590 }
591 shared.document = Some(reply.raw).filter(|raw| !raw.is_null());
592 shared.request_id = reply.provider_request_id;
593 let mut items = Vec::new();
594 while let Some(item) = shared.take() {
595 let failed = item.is_err();
596 items.push(item);
597 if failed {
598 break;
599 }
600 }
601 let outcome = match items.last() {
602 Some(Err(error)) => Err(error.clone()),
603 _ => shared.conclude(&reply.provider),
604 };
605 #[cfg(any(test, feature = "test-utils"))]
606 if let Err(error) = &outcome
607 && !matches!(items.last(), Some(Err(_)))
608 {
609 items.push(Err(error.clone()));
610 }
611 Decoded {
612 #[cfg(any(test, feature = "test-utils"))]
613 items,
614 outcome,
615 }
616}
617
618#[cfg(any(test, feature = "test-utils"))]
621pub(crate) fn relay_frames<W>(
622 wire: &W,
623 frames: impl IntoIterator<Item = W::Frame>,
624) -> crate::streaming::StreamEvents
625where
626 W: Wire<Op = crate::operation::Completion>,
627{
628 use crate::error::ErrorReport;
629 use crate::streaming::Relayed;
630
631 let provider = wire.describe().name.to_owned();
632 let shared = Mutex::new(Shared::new(crate::operation::Turn::relayed(
633 provider.clone(),
634 )));
635 let fed = feed(
636 &mut wire.decoder(),
637 Some(wire.reassembler()),
638 &shared,
639 frames,
640 );
641 let decoded = settle(
642 shared,
643 fed,
644 crate::wire::Reply {
645 provider,
646 raw: serde_json::Value::Null,
647 provider_request_id: None,
648 },
649 );
650 let origin = decoded
651 .outcome
652 .as_ref()
653 .map(|response| response.origin.clone())
654 .ok();
655 let relayed: Vec<Result<Relayed, ErrorReport>> = origin
656 .map(|origin| Ok(Relayed::Origin(origin)))
657 .into_iter()
658 .chain(decoded.items.into_iter().map(|item| match item {
659 Ok(item) => Ok(Relayed::Item(item)),
660 Err(error) => Err(ErrorReport::from(&error)),
661 }))
662 .chain(
664 decoded
665 .outcome
666 .ok()
667 .map(|response| Ok(Relayed::Done(Box::new(response)))),
668 )
669 .collect();
670 Box::pin(futures::stream::iter(relayed))
671}
672
673#[cfg(test)]
674impl Decoded<crate::operation::Completion> {
675 pub(crate) fn events(&self) -> Vec<&crate::streaming::StreamEvent> {
677 self.items
678 .iter()
679 .filter_map(|item| match item {
680 Ok(crate::streaming::Item::Event(event)) => Some(event),
681 _ => None,
682 })
683 .collect()
684 }
685
686 pub(crate) fn ended(&self) -> Vec<crate::message::AssistantContent> {
688 self.events()
689 .into_iter()
690 .filter_map(|event| match event {
691 crate::streaming::StreamEvent::End { content, .. } => Some(content.clone()),
692 _ => None,
693 })
694 .collect()
695 }
696}
697
698#[cfg(test)]
701macro_rules! decode_events {
702 ($decoder:expr, $provider:expr, $events:expr) => {
703 $crate::driver::decode_with(
704 $crate::operation::Turn::relayed($provider),
705 $provider,
706 |reply| {
707 let mut decoder = $decoder;
708 for event in $events {
709 if let $crate::wire::Flow::Ended(_) =
710 $crate::wire::Decoder::decode(&mut decoder, event, reply.out())?
711 {
712 return Ok(());
713 }
714 }
715 match $crate::wire::Decoder::eof(&mut decoder, reply.out())? {
716 $crate::wire::Flow::Ended(_) => Ok(()),
717 $crate::wire::Flow::More => Err($crate::error::ProviderError::Truncated),
718 }
719 },
720 )
721 };
722}
723#[cfg(test)]
724pub(crate) use decode_events;
725
726#[cfg(test)]
731macro_rules! feed_frames {
732 ($decoder:expr, $provider:expr, $frames:expr) => {
733 $crate::driver::decode_with(
734 $crate::operation::Turn::relayed($provider),
735 $provider,
736 |reply| {
737 let mut decoder = $decoder;
738 reply.feed(
739 &mut decoder,
740 None::<$crate::wire::document::Unreassembled>,
741 $frames,
742 )
743 },
744 )
745 };
746 ($decoder:expr, $reassembler:expr, $provider:expr, $frames:expr) => {
747 $crate::driver::decode_with(
748 $crate::operation::Turn::relayed($provider),
749 $provider,
750 |reply| {
751 let mut decoder = $decoder;
752 reply.feed(&mut decoder, Some($reassembler), $frames)
753 },
754 )
755 };
756}
757#[cfg(test)]
758pub(crate) use feed_frames;
759
760#[cfg(test)]
762pub(crate) struct Replying<'id, Op: Operation>(&'id Mutex<Shared<Op>>);
763
764#[cfg(test)]
765impl<'id, Op: Operation> Replying<'id, Op> {
766 pub(crate) fn out(&self) -> Out<'id, Op> {
768 Out::new(self.0)
769 }
770
771 pub(crate) fn feed<F, D: Decoder<'id, Op, F>, R: Reassemble<F>>(
773 &self,
774 decoder: &mut D,
775 reassembler: Option<R>,
776 frames: impl IntoIterator<Item = F>,
777 ) -> Result<(), ProviderError> {
778 feed(decoder, reassembler, self.0, frames)
779 }
780}
781
782#[cfg(test)]
786pub(crate) fn decode_with<Op: Operation>(
787 fold: Op::Fold,
788 provider: &str,
789 run: impl for<'id> FnOnce(Replying<'id, Op>) -> Result<(), ProviderError>,
790) -> Decoded<Op> {
791 let shared = Mutex::new(Shared::new(fold));
792 let fed = run(Replying(&shared));
793 settle(
794 shared,
795 fed,
796 crate::wire::Reply {
797 provider: provider.to_owned(),
798 raw: serde_json::Value::Null,
799 provider_request_id: None,
800 },
801 )
802}
803
804fn fail<Op: Operation>(
806 reply: &Mutex<Shared<Op>>,
807 slot: Option<&AdapterSlot>,
808 error: ProviderError,
809) {
810 if let Some(slot) = slot {
811 slot.fail(&error);
812 }
813 lock(reply).items.push_back(Err(error));
814}
815
816pub(crate) fn lock<T>(mutex: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
817 mutex
818 .lock()
819 .unwrap_or_else(std::sync::PoisonError::into_inner)
820}
821
822pub(crate) async fn follow_cursors<P, Fut>(
826 provider: &str,
827 operation: &str,
828 mut page: impl FnMut(Option<String>) -> Fut,
829) -> Result<Vec<P>, ProviderError>
830where
831 Fut: Future<Output = Result<(P, Option<String>), ProviderError>>,
832{
833 let mut pages = Vec::new();
834 let mut cursor = None;
835 loop {
836 let (read, next) = page(cursor.clone()).await?;
837 pages.push(read);
838 let Some(next) = next else { break };
841 if cursor.as_deref() == Some(next.as_str()) {
842 tracing::warn!(
845 provider,
846 operation,
847 pages = pages.len(),
848 "listing repeated its pagination cursor; returning the pages fetched so far"
849 );
850 break;
851 }
852 if pages.len() >= MAX_CONTINUATION_PAGES {
853 tracing::warn!(
854 provider,
855 operation,
856 pages = pages.len(),
857 "listing hit its page ceiling with a cursor still advancing; returning the pages \
858 fetched so far"
859 );
860 break;
861 }
862 cursor = Some(next);
863 }
864 Ok(pages)
865}
866
867pub fn warn_unmodeled(kind: &str, payload: &impl serde::Serialize) {
870 tracing::warn!(
871 kind,
872 payload_bytes = unknown_payload_bytes(payload),
873 "skipping unmodeled wire payload"
874 );
875}
876
877fn unknown_payload_bytes(value: &impl serde::Serialize) -> u64 {
880 struct CountingWriter(u64);
883
884 impl std::io::Write for CountingWriter {
885 fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
886 self.0 += buf.len() as u64;
887 Ok(buf.len())
888 }
889
890 fn flush(&mut self) -> std::io::Result<()> {
891 Ok(())
892 }
893 }
894
895 let mut counter = CountingWriter(0);
896 let _ = serde_json::to_writer(&mut counter, value);
898 counter.0
899}
900
901const MAX_CONTINUATION_PAGES: usize = 1000;
904
905pub(crate) fn record_request_id(span: &tracing::Span, request_id: Option<&str>) {
907 if let Some(request_id) = request_id
908 && !span.is_disabled()
909 {
910 span.record(crate::telemetry::PROVIDER_REQUEST_ID_FIELD, request_id);
911 }
912}
913
914#[cfg(test)]
915pub(crate) mod tests;