rig_core/providers/openai/wire/
route.rs1use serde::{Deserialize, Serialize};
11
12use crate::completion::CompletionRequest;
13use crate::error::{EncodeError, ProviderError};
14use crate::operation::Completion;
15use crate::providers::openai::responses_api::streaming::{ResponsesDecoder, ResponsesEvent};
16use crate::providers::openai::responses_api::wire::Responses;
17use crate::providers::openai::responses_api::{
18 ResponsesToolDefinition, SystemInstructionsPlacement,
19};
20use crate::wire::document::Reassemble;
21use crate::wire::{Decoder, Descriptor, Encoded, Flow, Mode, Out, Wire, WireEvent, WireFrame};
22
23use super::OpenAIConfig;
24use super::chat::{Chat, ChatDecoder, ChatEvent};
25
26macro_rules! on_route {
28 ($chosen:expr, $wire:ident => $ask:expr) => {
29 match $chosen {
30 Self::Chat($wire) => $ask,
31 Self::Responses($wire) => $ask,
32 }
33 };
34}
35
36#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
38pub enum Route {
39 Chat,
41 Responses,
43}
44
45#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
48pub enum OpenAiWire {
49 Chat(Chat),
51 Responses(Responses),
53}
54
55impl OpenAiWire {
56 pub fn new(provider: OpenAIConfig, model: impl Into<String>) -> Self {
59 let model = model.into();
60 let route = provider.route.unwrap_or_else(|| {
61 provider
62 .dialect
63 .quirks
64 .hooks
65 .and_then(|hooks| hooks.model_route)
66 .map_or_else(|| provider.completion_route(), |route| route(&model))
67 });
68 match route {
69 Route::Chat => Self::Chat(Chat::new(provider, model)),
70 Route::Responses => Self::Responses(Responses::new(provider, model)),
71 }
72 }
73
74 pub(crate) fn encode_with_headers(
77 &self,
78 request: CompletionRequest,
79 mode: Mode,
80 headers: impl FnOnce(
81 &OpenAIConfig,
82 &CompletionRequest,
83 http::request::Builder,
84 ) -> http::request::Builder,
85 ) -> Result<Encoded, EncodeError> {
86 on_route!(self, wire => wire.encode_with_headers(request, mode, headers))
87 }
88
89 pub fn provider(&self) -> &OpenAIConfig {
91 on_route!(self, wire => &wire.provider)
92 }
93
94 pub fn with_strict_tools(self) -> Self {
97 self.on_chat(Chat::with_strict_tools)
98 .on_responses(Responses::with_strict_tools)
99 }
100
101 pub fn with_tool_result_array_content(self) -> Self {
105 self.on_chat(Chat::with_tool_result_array_content)
106 }
107
108 pub fn with_tool(self, tool: impl Into<ResponsesToolDefinition>) -> Self {
111 self.on_responses(|wire| wire.with_tool(tool))
112 }
113
114 pub fn with_tools<I, Tool>(self, tools: I) -> Self
117 where
118 I: IntoIterator<Item = Tool>,
119 Tool: Into<ResponsesToolDefinition>,
120 {
121 self.on_responses(|wire| wire.with_tools(tools))
122 }
123
124 pub fn with_system_instructions_placement(
128 self,
129 placement: SystemInstructionsPlacement,
130 ) -> Self {
131 self.on_responses(|wire| wire.with_system_instructions_placement(placement))
132 }
133
134 pub fn with_system_instructions_as_messages(self) -> Self {
137 self.on_responses(Responses::with_system_instructions_as_messages)
138 }
139
140 fn on_chat(self, option: impl FnOnce(Chat) -> Chat) -> Self {
142 match self {
143 Self::Chat(wire) => Self::Chat(option(wire)),
144 responses => responses,
145 }
146 }
147
148 fn on_responses(self, option: impl FnOnce(Responses) -> Responses) -> Self {
150 match self {
151 Self::Responses(wire) => Self::Responses(option(wire)),
152 chat => chat,
153 }
154 }
155}
156
157impl From<Chat> for OpenAiWire {
158 fn from(wire: Chat) -> Self {
159 Self::Chat(wire)
160 }
161}
162
163impl From<Responses> for OpenAiWire {
164 fn from(wire: Responses) -> Self {
165 Self::Responses(wire)
166 }
167}
168
169impl Wire for OpenAiWire {
170 type Op = Completion;
171 type Payload = crate::wire::Encoded;
172 type Frame = crate::wire::WireFrame;
173 type Decoder<'id> = OpenAiDecoder;
174 type Reassembler = OpenAiReassembler;
175
176 fn describe(&self) -> Descriptor<'_> {
177 on_route!(self, wire => wire.describe())
178 }
179
180 fn encode(&self, request: CompletionRequest, mode: Mode) -> Result<Encoded, EncodeError> {
181 on_route!(self, wire => wire.encode(request, mode))
182 }
183
184 fn decoder<'id>(&self) -> Self::Decoder<'id> {
185 match self {
186 Self::Chat(wire) => OpenAiDecoder::Chat(wire.decoder()),
187 Self::Responses(wire) => OpenAiDecoder::Responses(wire.decoder()),
188 }
189 }
190
191 fn reassembler(&self) -> Self::Reassembler {
192 match self {
193 Self::Chat(wire) => OpenAiReassembler::Chat(wire.reassembler()),
194 Self::Responses(wire) => OpenAiReassembler::Responses(wire.reassembler()),
195 }
196 }
197}
198
199pub enum OpenAiReassembler {
201 Chat(<Chat as Wire>::Reassembler),
203 Responses(<Responses as Wire>::Reassembler),
205}
206
207impl Default for OpenAiReassembler {
209 fn default() -> Self {
210 Self::Chat(Default::default())
211 }
212}
213
214impl crate::wire::document::Serves<crate::operation::Completion> for OpenAiReassembler {}
215
216impl Reassemble<WireFrame> for OpenAiReassembler {
217 fn absorb(&mut self, frame: &WireFrame) {
218 on_route!(self, document => document.absorb(frame));
219 }
220
221 fn finish(self) -> serde_json::Value {
222 on_route!(self, document => document.finish())
223 }
224}
225
226pub enum OpenAiEvent {
228 Chat(ChatEvent),
230 Responses(ResponsesEvent),
232}
233
234pub enum OpenAiDecoder {
236 Chat(ChatDecoder),
238 Responses(ResponsesDecoder),
240}
241
242impl<'id> Decoder<'id, Completion> for OpenAiDecoder {
243 type Event = OpenAiEvent;
244
245 fn classify(&self, frame: WireFrame) -> WireEvent<OpenAiEvent> {
246 match self {
247 Self::Chat(decoder) => decoder.classify(frame).map(OpenAiEvent::Chat),
248 Self::Responses(decoder) => decoder.classify(frame).map(OpenAiEvent::Responses),
249 }
250 }
251
252 fn decode(
253 &mut self,
254 event: OpenAiEvent,
255 out: Out<'id, Completion>,
256 ) -> Result<Flow, ProviderError> {
257 match (self, event) {
258 (Self::Chat(decoder), OpenAiEvent::Chat(event)) => decoder.decode(event, out),
259 (Self::Responses(decoder), OpenAiEvent::Responses(event)) => decoder.decode(event, out),
260 (Self::Chat(_), OpenAiEvent::Responses(_))
263 | (Self::Responses(_), OpenAiEvent::Chat(_)) => Ok(Flow::More),
264 }
265 }
266
267 fn eof(&mut self, out: Out<'id, Completion>) -> Result<Flow, ProviderError> {
268 on_route!(self, decoder => decoder.eof(out))
269 }
270}
271
272#[cfg(test)]
273mod tests;