1use crate::protocol::operation::{
6 ContentGenerationKind, HttpMethod, Operation, OperationKey, OperationKind, Provider,
7};
8
9#[derive(Debug, Clone, PartialEq, Eq, gproxy_protocol_macros::WireBuilder)]
11#[non_exhaustive]
12pub struct RequestTarget {
13 pub method: HttpMethod,
14 pub path: String,
15 pub query: Option<String>,
17}
18
19impl RequestTarget {
20 fn get(path: impl Into<String>) -> Self {
21 Self {
22 method: HttpMethod::Get,
23 path: path.into(),
24 query: None,
25 }
26 }
27
28 fn post(path: impl Into<String>) -> Self {
29 Self {
30 method: HttpMethod::Post,
31 path: path.into(),
32 query: None,
33 }
34 }
35}
36
37#[derive(Debug, Clone, PartialEq, Eq)]
38#[non_exhaustive]
39pub enum EndpointError {
40 InconsistentOperationKey(OperationKey),
41 StreamMismatch {
42 operation: Operation,
43 stream: bool,
44 },
45 UnsupportedOperation {
46 operation: Operation,
47 provider: Provider,
48 },
49 MissingModel {
50 operation: Operation,
51 provider: Provider,
52 },
53}
54
55impl std::fmt::Display for EndpointError {
56 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
57 match self {
58 Self::InconsistentOperationKey(key) => {
59 write!(formatter, "inconsistent operation key: {key:?}")
60 }
61 Self::StreamMismatch { operation, stream } => write!(
62 formatter,
63 "operation {operation:?} is incompatible with stream={stream}"
64 ),
65 Self::UnsupportedOperation {
66 operation,
67 provider,
68 } => write!(
69 formatter,
70 "operation {operation:?} is unsupported by {provider:?}"
71 ),
72 Self::MissingModel {
73 operation,
74 provider,
75 } => write!(
76 formatter,
77 "operation {operation:?} for {provider:?} requires a non-empty raw model id"
78 ),
79 }
80 }
81}
82
83impl std::error::Error for EndpointError {}
84
85pub fn request_target(
89 target: OperationKey,
90 model: &str,
91 stream: bool,
92) -> Result<RequestTarget, EndpointError> {
93 if !target.is_consistent() {
94 return Err(EndpointError::InconsistentOperationKey(target));
95 }
96 use Provider as P;
97 let provider = match target.kind() {
98 OperationKind::ContentGeneration(kind) => {
99 let operation_streams = target.operation() == Operation::StreamGenerateContent;
100 if operation_streams != stream {
101 return Err(EndpointError::StreamMismatch {
102 operation: target.operation(),
103 stream,
104 });
105 }
106 return content_target(target.operation(), kind, model, stream);
107 }
108 OperationKind::Provider(provider) => provider,
109 };
110 let request_target = match (target.operation(), provider) {
111 (Operation::ListModels, P::OpenAi | P::Claude) => RequestTarget::get("/v1/models"),
112 (Operation::ListModels, P::Gemini) => RequestTarget::get("/v1beta/models"),
113 (Operation::GetModel, P::OpenAi | P::Claude) => {
114 require_model(target.operation(), provider, model)?;
115 RequestTarget::get(format!("/v1/models/{}", encode_component(model)))
116 }
117 (Operation::GetModel, P::Gemini) => {
118 require_model(target.operation(), provider, model)?;
119 RequestTarget::get(format!("/v1beta/models/{}", encode_component(model)))
120 }
121 (Operation::CountTokens, P::OpenAi) => RequestTarget::post("/v1/responses/input_tokens"),
122 (Operation::CountTokens, P::Claude) => RequestTarget::post("/v1/messages/count_tokens"),
123 (Operation::CountTokens, P::Gemini) => {
124 require_model(target.operation(), provider, model)?;
125 RequestTarget::post(format!(
126 "/v1beta/models/{}:countTokens",
127 encode_component(model)
128 ))
129 }
130 (Operation::CreateEmbedding, P::OpenAi) => RequestTarget::post("/v1/embeddings"),
131 (Operation::CreateEmbedding, P::Gemini) => {
133 require_model(target.operation(), provider, model)?;
134 RequestTarget::post(format!(
135 "/v1beta/models/{}:embedContent",
136 encode_component(model)
137 ))
138 }
139 (Operation::CreateImage, P::OpenAi) => RequestTarget::post("/v1/images/generations"),
140 (Operation::EditImage, P::OpenAi) => RequestTarget::post("/v1/images/edits"),
141 (Operation::CompactContent, P::OpenAi) => RequestTarget::post("/v1/responses/compact"),
142 (Operation::CreateConversation, P::OpenAi) => RequestTarget::post("/v1/conversations"),
143 (Operation::ConnectRealtime, P::OpenAi) => {
144 require_model(target.operation(), provider, model)?;
145 RequestTarget {
146 method: HttpMethod::Get,
147 path: "/v1/realtime".to_owned(),
148 query: Some(format!("model={}", encode_component(model))),
149 }
150 }
151 (operation, provider) => {
152 return Err(EndpointError::UnsupportedOperation {
153 operation,
154 provider,
155 });
156 }
157 };
158 Ok(request_target)
159}
160
161fn content_target(
163 operation: Operation,
164 kind: ContentGenerationKind,
165 model: &str,
166 stream: bool,
167) -> Result<RequestTarget, EndpointError> {
168 use ContentGenerationKind as K;
169 let target = match kind {
170 K::OpenAiChatCompletions => RequestTarget::post("/v1/chat/completions"),
171 K::OpenAiResponses => RequestTarget::post("/v1/responses"),
172 K::OpenAiResponsesWebSocket if stream => RequestTarget::get("/v1/responses"),
173 K::OpenAiResponsesWebSocket => {
174 return Err(EndpointError::StreamMismatch { operation, stream });
175 }
176 K::ClaudeMessages => RequestTarget::post("/v1/messages"),
177 K::GeminiGenerateContent => {
178 require_model(operation, Provider::Gemini, model)?;
179 let verb = if stream {
180 "streamGenerateContent"
181 } else {
182 "generateContent"
183 };
184 RequestTarget {
185 method: HttpMethod::Post,
186 path: format!("/v1beta/models/{}:{verb}", encode_component(model)),
187 query: stream.then(|| "alt=sse".to_owned()),
188 }
189 }
190 };
191 Ok(target)
192}
193
194fn require_model(
195 operation: Operation,
196 provider: Provider,
197 model: &str,
198) -> Result<(), EndpointError> {
199 if model.is_empty() {
200 Err(EndpointError::MissingModel {
201 operation,
202 provider,
203 })
204 } else {
205 Ok(())
206 }
207}
208
209fn encode_component(value: &str) -> String {
212 const HEX: &[u8; 16] = b"0123456789ABCDEF";
213 let mut encoded = String::with_capacity(value.len());
214 for byte in value.bytes() {
215 if byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'.' | b'_' | b'~') {
216 encoded.push(char::from(byte));
217 } else {
218 encoded.push('%');
219 encoded.push(char::from(HEX[usize::from(byte >> 4)]));
220 encoded.push(char::from(HEX[usize::from(byte & 0x0f)]));
221 }
222 }
223 encoded
224}
225
226#[cfg(test)]
227mod tests {
228 use super::*;
229
230 #[test]
231 fn provider_endpoint_support_matrix_is_explicit() {
232 use Operation as O;
233 use Provider as P;
234 let rows = [
235 (O::ListModels, [true, true, true]),
236 (O::GetModel, [true, true, true]),
237 (O::CountTokens, [true, true, true]),
238 (O::CreateEmbedding, [true, false, true]),
239 (O::CreateImage, [true, false, false]),
240 (O::EditImage, [true, false, false]),
241 (O::CompactContent, [true, false, false]),
242 (O::CreateConversation, [true, false, false]),
243 (O::ConnectRealtime, [true, false, false]),
244 ];
245 for (operation, supported) in rows {
246 for (provider, expected) in [P::OpenAi, P::Claude, P::Gemini].into_iter().zip(supported)
247 {
248 let result =
249 request_target(OperationKey::provider(operation, provider), "model", false);
250 assert_eq!(result.is_ok(), expected, "{operation:?} / {provider:?}");
251 }
252 }
253 }
254
255 #[test]
256 fn content_endpoint_support_matrix_is_explicit() {
257 use ContentGenerationKind as K;
258 for kind in [
259 K::OpenAiResponses,
260 K::OpenAiResponsesWebSocket,
261 K::OpenAiChatCompletions,
262 K::ClaudeMessages,
263 K::GeminiGenerateContent,
264 ] {
265 for (operation, stream) in [
266 (Operation::GenerateContent, false),
267 (Operation::StreamGenerateContent, true),
268 ] {
269 let result = request_target(
270 OperationKey::content_generation(operation, kind),
271 "model",
272 stream,
273 );
274 assert_eq!(
275 result.is_ok(),
276 kind != K::OpenAiResponsesWebSocket || stream,
277 "{operation:?} / {kind:?}"
278 );
279 }
280 }
281 }
282
283 #[test]
284 fn rejects_inconsistent_keys_and_stream_flags() {
285 let inconsistent = OperationKey::new_unchecked(
286 Operation::GenerateContent,
287 OperationKind::Provider(Provider::OpenAi),
288 );
289 assert!(matches!(
290 request_target(inconsistent, "model", false),
291 Err(EndpointError::InconsistentOperationKey(_))
292 ));
293
294 let key = OperationKey::content_generation(
295 Operation::StreamGenerateContent,
296 ContentGenerationKind::OpenAiChatCompletions,
297 );
298 assert!(matches!(
299 request_target(key, "model", false),
300 Err(EndpointError::StreamMismatch { .. })
301 ));
302 }
303
304 #[test]
305 fn model_is_a_raw_encoded_component() {
306 let key = OperationKey::provider(Operation::GetModel, Provider::Gemini);
307 let target = request_target(key, "org/model ?", false).unwrap();
308 assert_eq!(target.path, "/v1beta/models/org%2Fmodel%20%3F");
309
310 let key = OperationKey::provider(Operation::ConnectRealtime, Provider::OpenAi);
311 let target = request_target(key, "gpt/a b", false).unwrap();
312 assert_eq!(target.query.as_deref(), Some("model=gpt%2Fa%20b"));
313 }
314}