Skip to main content

rig_core/driver/
local.rs

1//! The pass-through wire: the payload is the request and each frame is one
2//! step of the reply. It is what a local runtime or an in-process model
3//! speaks, so its transport is the runtime itself.
4//!
5//! ```
6//! use rig_core::driver::Local;
7//! use rig_core::operation::Embedding;
8//! use rig_core::wire::{Capabilities, Wire};
9//!
10//! let wire = Local::<Embedding>::new("local").with_capabilities(Capabilities::embedding(8, 3));
11//! assert_eq!(wire.describe().name, "local");
12//! assert_eq!(wire.describe().capabilities.ndims, 3);
13//! ```
14//!
15//! A completion's events come only from its writer's part handles, so there
16//! is no local completion wire:
17//!
18//! ```compile_fail
19//! use rig_core::driver::Local;
20//! use rig_core::operation::Completion;
21//!
22//! let wire = Local::<Completion>::new("scripted");
23//! ```
24
25use std::fmt;
26use std::marker::PhantomData;
27
28use crate::error::{EncodeError, ProviderError};
29use crate::streaming::UnknownPayload;
30use crate::wire::{
31    Capabilities, Decoder, Descriptor, Flow, Free, Mode, Operation, Out, Wire, WireEvent,
32};
33
34/// A wire whose payload is the request and whose frames are the reply's
35/// steps, for an operation whose events the runtime builds itself.
36///
37/// The transport is the runtime: it implements `Transport<Local<Op>>`,
38/// takes the request and delivers the reply's [`Step`]s, ending with the
39/// operation's end. A runtime that stops without it leaves the reply
40/// truncated. A failure that ends the reply is the transport's own `Err`.
41pub struct Local<Op> {
42    name: String,
43    id: Option<String>,
44    capabilities: Capabilities,
45    op: PhantomData<fn() -> Op>,
46}
47
48/// One step of a local runtime's reply.
49pub enum Step<Op: Operation> {
50    /// An event of the reply.
51    Event(Op::Event),
52    /// A payload the runtime does not model.
53    Unknown(UnknownPayload),
54    /// The runtime's end of the reply.
55    End(Op::End),
56}
57
58impl<Op: Operation<Emit = Free>> Local<Op> {
59    /// A wire named `name` (as records and telemetry name it), addressing
60    /// no model id, with default capabilities.
61    pub fn new(name: impl Into<String>) -> Self {
62        Self {
63            name: name.into(),
64            id: None,
65            capabilities: Capabilities::default(),
66            op: PhantomData,
67        }
68    }
69
70    /// What a runtime accounts for, such as an embedding width.
71    pub fn with_capabilities(mut self, capabilities: Capabilities) -> Self {
72        self.capabilities = capabilities;
73        self
74    }
75
76    /// The model id this wire addresses.
77    pub fn with_id(mut self, id: impl Into<String>) -> Self {
78        self.id = Some(id.into());
79        self
80    }
81}
82
83impl<Op> Clone for Local<Op> {
84    fn clone(&self) -> Self {
85        Self {
86            name: self.name.clone(),
87            id: self.id.clone(),
88            capabilities: self.capabilities,
89            op: PhantomData,
90        }
91    }
92}
93
94impl<Op> fmt::Debug for Local<Op> {
95    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
96        f.debug_struct("Local")
97            .field("name", &self.name)
98            .field("id", &self.id)
99            .field("capabilities", &self.capabilities)
100            .finish()
101    }
102}
103
104impl<Op: Operation<Emit = Free>> Wire for Local<Op> {
105    type Op = Op;
106    type Payload = Op::Request;
107    type Frame = Step<Op>;
108    type Decoder<'id> = Steps;
109    type Reassembler = crate::wire::document::Unreassembled;
110
111    fn describe(&self) -> Descriptor<'_> {
112        Descriptor::new(&self.name)
113            .model(self.id.as_deref())
114            .capabilities(self.capabilities)
115    }
116
117    fn encode(&self, request: Op::Request, _mode: Mode) -> Result<Op::Request, EncodeError> {
118        Ok(request)
119    }
120
121    fn decoder<'id>(&self) -> Self::Decoder<'id> {
122        Steps
123    }
124}
125
126/// The decoder of a [`Local`] wire: every step is written as it is.
127#[derive(Debug, Default)]
128pub struct Steps;
129
130impl<'id, Op: Operation<Emit = Free>> Decoder<'id, Op, Step<Op>> for Steps {
131    type Event = Step<Op>;
132
133    fn classify(&self, step: Step<Op>) -> WireEvent<Step<Op>> {
134        WireEvent::Known(step)
135    }
136
137    fn decode(&mut self, step: Step<Op>, mut out: Out<'id, Op>) -> Result<Flow, ProviderError> {
138        Ok(match step {
139            Step::Event(event) => {
140                out.event(event);
141                Flow::More
142            }
143            Step::Unknown(payload) => {
144                out.unknown(payload);
145                Flow::More
146            }
147            Step::End(end) => out.end(end),
148        })
149    }
150}