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::{Decoder, Descriptor, Encoded, Flow, Mode, Out, Wire, WireEvent, WireFrame};
21
22use super::OpenAIConfig;
23use super::chat::{Chat, ChatDecoder, ChatEvent};
24
25macro_rules! on_route {
27 ($chosen:expr, $wire:ident => $ask:expr) => {
28 match $chosen {
29 Self::Chat($wire) => $ask,
30 Self::Responses($wire) => $ask,
31 }
32 };
33}
34
35#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
37pub enum Route {
38 Chat,
40 Responses,
42}
43
44#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
47pub enum OpenAiWire {
48 Chat(Chat),
50 Responses(Responses),
52}
53
54impl OpenAiWire {
55 pub fn new(provider: OpenAIConfig, model: impl Into<String>) -> Self {
58 let model = model.into();
59 let route = provider.route.unwrap_or_else(|| {
60 provider
61 .dialect
62 .quirks
63 .hooks
64 .and_then(|hooks| hooks.model_route)
65 .map_or_else(|| provider.completion_route(), |route| route(&model))
66 });
67 match route {
68 Route::Chat => Self::Chat(Chat::new(provider, model)),
69 Route::Responses => Self::Responses(Responses::new(provider, model)),
70 }
71 }
72
73 pub(crate) fn encode_with_headers(
76 &self,
77 request: CompletionRequest,
78 mode: Mode,
79 headers: impl FnOnce(
80 &OpenAIConfig,
81 &CompletionRequest,
82 http::request::Builder,
83 ) -> http::request::Builder,
84 ) -> Result<Encoded, EncodeError> {
85 on_route!(self, wire => wire.encode_with_headers(request, mode, headers))
86 }
87
88 pub fn provider(&self) -> &OpenAIConfig {
90 on_route!(self, wire => &wire.provider)
91 }
92
93 pub fn with_strict_tools(self) -> Self {
96 self.on_chat(Chat::with_strict_tools)
97 .on_responses(Responses::with_strict_tools)
98 }
99
100 pub fn with_tool_result_array_content(self) -> Self {
104 self.on_chat(Chat::with_tool_result_array_content)
105 }
106
107 pub fn with_prompt_caching(self) -> Self {
110 self.on_chat(Chat::with_prompt_caching)
111 }
112
113 pub fn with_tool(self, tool: impl Into<ResponsesToolDefinition>) -> Self {
116 self.on_responses(|wire| wire.with_tool(tool))
117 }
118
119 pub fn with_tools<I, Tool>(self, tools: I) -> Self
122 where
123 I: IntoIterator<Item = Tool>,
124 Tool: Into<ResponsesToolDefinition>,
125 {
126 self.on_responses(|wire| wire.with_tools(tools))
127 }
128
129 pub fn with_system_instructions_placement(
133 self,
134 placement: SystemInstructionsPlacement,
135 ) -> Self {
136 self.on_responses(|wire| wire.with_system_instructions_placement(placement))
137 }
138
139 pub fn with_system_instructions_as_messages(self) -> Self {
142 self.on_responses(Responses::with_system_instructions_as_messages)
143 }
144
145 fn on_chat(self, option: impl FnOnce(Chat) -> Chat) -> Self {
147 match self {
148 Self::Chat(wire) => Self::Chat(option(wire)),
149 responses => responses,
150 }
151 }
152
153 fn on_responses(self, option: impl FnOnce(Responses) -> Responses) -> Self {
155 match self {
156 Self::Responses(wire) => Self::Responses(option(wire)),
157 chat => chat,
158 }
159 }
160}
161
162impl From<Chat> for OpenAiWire {
163 fn from(wire: Chat) -> Self {
164 Self::Chat(wire)
165 }
166}
167
168impl From<Responses> for OpenAiWire {
169 fn from(wire: Responses) -> Self {
170 Self::Responses(wire)
171 }
172}
173
174impl Wire for OpenAiWire {
175 type Op = Completion;
176 type Payload = crate::wire::Encoded;
177 type Frame = crate::wire::WireFrame;
178 type Decoder<'id> = OpenAiDecoder<'id>;
179
180 fn describe(&self) -> Descriptor<'_> {
181 on_route!(self, wire => wire.describe())
182 }
183
184 fn encode(&self, request: CompletionRequest, mode: Mode) -> Result<Encoded, EncodeError> {
185 on_route!(self, wire => wire.encode(request, mode))
186 }
187
188 fn decoder<'id>(&self) -> OpenAiDecoder<'id> {
189 match self {
190 Self::Chat(wire) => OpenAiDecoder::Chat(wire.decoder()),
191 Self::Responses(wire) => OpenAiDecoder::Responses(wire.decoder()),
192 }
193 }
194}
195
196pub enum OpenAiEvent {
198 Chat(ChatEvent),
200 Responses(ResponsesEvent),
202}
203
204pub enum OpenAiDecoder<'id> {
206 Chat(ChatDecoder<'id>),
208 Responses(ResponsesDecoder<'id>),
210}
211
212impl<'id> Decoder<'id, Completion> for OpenAiDecoder<'id> {
213 type Event = OpenAiEvent;
214
215 fn classify(&self, frame: WireFrame) -> WireEvent<OpenAiEvent> {
216 match self {
217 Self::Chat(decoder) => decoder.classify(frame).map(OpenAiEvent::Chat),
218 Self::Responses(decoder) => decoder.classify(frame).map(OpenAiEvent::Responses),
219 }
220 }
221
222 fn decode(
223 &mut self,
224 event: OpenAiEvent,
225 out: Out<'id, Completion>,
226 ) -> Result<Flow, ProviderError> {
227 match (self, event) {
228 (Self::Chat(decoder), OpenAiEvent::Chat(event)) => decoder.decode(event, out),
229 (Self::Responses(decoder), OpenAiEvent::Responses(event)) => decoder.decode(event, out),
230 (Self::Chat(_), OpenAiEvent::Responses(_))
233 | (Self::Responses(_), OpenAiEvent::Chat(_)) => Ok(Flow::More),
234 }
235 }
236
237 fn eof(&mut self, out: Out<'id, Completion>) -> Result<Flow, ProviderError> {
238 on_route!(self, decoder => decoder.eof(out))
239 }
240}
241
242#[cfg(test)]
243mod tests;