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::CreateSpeech, P::OpenAi) => RequestTarget::post("/v1/audio/speech"),
132 (Operation::CreateTranscription, P::OpenAi) => {
133 RequestTarget::post("/v1/audio/transcriptions")
134 }
135 (Operation::CreateTranslation, P::OpenAi) => RequestTarget::post("/v1/audio/translations"),
136 (Operation::Rerank, P::OpenAi) => RequestTarget::post("/v1/rerank"),
137 (Operation::CreateEmbedding, P::Gemini) => {
139 require_model(target.operation(), provider, model)?;
140 RequestTarget::post(format!(
141 "/v1beta/models/{}:embedContent",
142 encode_component(model)
143 ))
144 }
145 (Operation::CreateImage, P::OpenAi) => RequestTarget::post("/v1/images/generations"),
146 (Operation::CreateImage, P::Gemini) => {
149 require_model(target.operation(), provider, model)?;
150 RequestTarget::post(format!(
151 "/v1beta/models/{}:predict",
152 encode_component(model)
153 ))
154 }
155 (Operation::EditImage, P::OpenAi) => RequestTarget::post("/v1/images/edits"),
156 (Operation::WebSearch, P::OpenAi) => RequestTarget::post("/v1/alpha/search"),
157 (Operation::CompactContent, P::OpenAi) => RequestTarget::post("/v1/responses/compact"),
158 (Operation::CreateConversation, P::OpenAi) => RequestTarget::post("/v1/conversations"),
159 (Operation::CreateVideo, P::OpenAi) => RequestTarget::post("/v1/videos"),
162 (Operation::ListVideos, P::OpenAi) => RequestTarget::get("/v1/videos"),
163 (Operation::RetrieveVideo, P::OpenAi) => {
164 require_model(target.operation(), provider, model)?;
165 RequestTarget::get(format!("/v1/videos/{}", encode_component(model)))
166 }
167 (Operation::DeleteVideo, P::OpenAi) => {
168 require_model(target.operation(), provider, model)?;
169 RequestTarget {
170 method: HttpMethod::Delete,
171 path: format!("/v1/videos/{}", encode_component(model)),
172 query: None,
173 }
174 }
175 (Operation::DownloadVideoContent, P::OpenAi) => {
176 require_model(target.operation(), provider, model)?;
177 RequestTarget::get(format!("/v1/videos/{}/content", encode_component(model)))
178 }
179 (Operation::RemixVideo, P::OpenAi) => {
180 require_model(target.operation(), provider, model)?;
181 RequestTarget::post(format!("/v1/videos/{}/remix", encode_component(model)))
182 }
183 (Operation::CreateVideoCharacter, P::OpenAi) => {
184 RequestTarget::post("/v1/videos/characters")
185 }
186 (Operation::GetVideoCharacter, P::OpenAi) => {
187 require_model(target.operation(), provider, model)?;
188 RequestTarget::get(format!("/v1/videos/characters/{}", encode_component(model)))
189 }
190 (Operation::EditVideo, P::OpenAi) => RequestTarget::post("/v1/videos/edits"),
191 (Operation::ExtendVideo, P::OpenAi) => RequestTarget::post("/v1/videos/extensions"),
192 (Operation::CreateVideo, P::Gemini) => {
193 require_model(target.operation(), provider, model)?;
194 RequestTarget::post(format!(
195 "/v1beta/models/{}:predictLongRunning",
196 encode_component(model)
197 ))
198 }
199 (Operation::RetrieveVideo, P::Gemini) => {
202 require_model(target.operation(), provider, model)?;
203 RequestTarget::get(format!("/v1beta/{}", model.trim_start_matches('/')))
204 }
205 (Operation::DownloadVideoContent, P::Gemini) => {
206 require_model(target.operation(), provider, model)?;
207 RequestTarget {
208 method: HttpMethod::Get,
209 path: format!("/v1beta/files/{}:download", encode_component(model)),
210 query: Some("alt=media".to_owned()),
211 }
212 }
213 (Operation::CreateRealtimeCall, P::OpenAi) => RequestTarget::post("/v1/realtime/calls"),
214 (Operation::CreateFile, P::OpenAi | P::Claude) => RequestTarget::post("/v1/files"),
215 (Operation::ListFiles, P::OpenAi | P::Claude) => RequestTarget::get("/v1/files"),
216 (Operation::RetrieveFile, P::OpenAi | P::Claude) => {
217 require_model(target.operation(), provider, model)?;
218 RequestTarget::get(format!("/v1/files/{}", encode_component(model)))
219 }
220 (Operation::DeleteFile, P::OpenAi | P::Claude) => {
221 require_model(target.operation(), provider, model)?;
222 RequestTarget {
223 method: HttpMethod::Delete,
224 path: format!("/v1/files/{}", encode_component(model)),
225 query: None,
226 }
227 }
228 (Operation::DownloadFileContent, P::OpenAi | P::Claude) => {
229 require_model(target.operation(), provider, model)?;
230 RequestTarget::get(format!("/v1/files/{}/content", encode_component(model)))
231 }
232 (Operation::ConnectRealtime, P::OpenAi) => {
233 require_model(target.operation(), provider, model)?;
234 RequestTarget {
235 method: HttpMethod::Get,
236 path: "/v1/realtime".to_owned(),
237 query: Some(format!("model={}", encode_component(model))),
238 }
239 }
240 (operation, provider) => {
241 return Err(EndpointError::UnsupportedOperation {
242 operation,
243 provider,
244 });
245 }
246 };
247 Ok(request_target)
248}
249
250fn content_target(
252 operation: Operation,
253 kind: ContentGenerationKind,
254 model: &str,
255 stream: bool,
256) -> Result<RequestTarget, EndpointError> {
257 use ContentGenerationKind as K;
258 let target = match kind {
259 K::OpenAiChatCompletions => RequestTarget::post("/v1/chat/completions"),
260 K::OpenAiResponses => RequestTarget::post("/v1/responses"),
261 K::OpenAiResponsesWebSocket if stream => RequestTarget::get("/v1/responses"),
262 K::OpenAiResponsesWebSocket => {
263 return Err(EndpointError::StreamMismatch { operation, stream });
264 }
265 K::ClaudeMessages => RequestTarget::post("/v1/messages"),
266 K::GeminiGenerateContent => {
267 require_model(operation, Provider::Gemini, model)?;
268 let verb = if stream {
269 "streamGenerateContent"
270 } else {
271 "generateContent"
272 };
273 RequestTarget {
274 method: HttpMethod::Post,
275 path: format!("/v1beta/models/{}:{verb}", encode_component(model)),
276 query: stream.then(|| "alt=sse".to_owned()),
277 }
278 }
279 };
280 Ok(target)
281}
282
283fn require_model(
284 operation: Operation,
285 provider: Provider,
286 model: &str,
287) -> Result<(), EndpointError> {
288 if model.is_empty() {
289 Err(EndpointError::MissingModel {
290 operation,
291 provider,
292 })
293 } else {
294 Ok(())
295 }
296}
297
298fn encode_component(value: &str) -> String {
301 const HEX: &[u8; 16] = b"0123456789ABCDEF";
302 let mut encoded = String::with_capacity(value.len());
303 for byte in value.bytes() {
304 if byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'.' | b'_' | b'~') {
305 encoded.push(char::from(byte));
306 } else {
307 encoded.push('%');
308 encoded.push(char::from(HEX[usize::from(byte >> 4)]));
309 encoded.push(char::from(HEX[usize::from(byte & 0x0f)]));
310 }
311 }
312 encoded
313}
314
315#[cfg(test)]
316mod tests {
317 use super::*;
318
319 #[test]
320 fn provider_endpoint_support_matrix_is_explicit() {
321 use Operation as O;
322 use Provider as P;
323 let rows = [
324 (O::ListModels, [true, true, true]),
325 (O::GetModel, [true, true, true]),
326 (O::CountTokens, [true, true, true]),
327 (O::CreateEmbedding, [true, false, true]),
328 (O::CreateSpeech, [true, false, false]),
329 (O::CreateTranscription, [true, false, false]),
330 (O::CreateTranslation, [true, false, false]),
331 (O::Rerank, [true, false, false]),
332 (O::CreateImage, [true, false, true]),
333 (O::EditImage, [true, false, false]),
334 (O::WebSearch, [true, false, false]),
335 (O::CompactContent, [true, false, false]),
336 (O::CreateConversation, [true, false, false]),
337 (O::CreateRealtimeCall, [true, false, false]),
338 (O::ConnectRealtime, [true, false, false]),
339 (O::CreateVideo, [true, false, true]),
340 (O::RetrieveVideo, [true, false, true]),
341 (O::ListVideos, [true, false, false]),
342 (O::DeleteVideo, [true, false, false]),
343 (O::DownloadVideoContent, [true, false, true]),
344 (O::RemixVideo, [true, false, false]),
345 (O::CreateVideoCharacter, [true, false, false]),
346 (O::GetVideoCharacter, [true, false, false]),
347 (O::EditVideo, [true, false, false]),
348 (O::ExtendVideo, [true, false, false]),
349 (O::CreateFile, [true, true, false]),
350 (O::ListFiles, [true, true, false]),
351 (O::RetrieveFile, [true, true, false]),
352 (O::DeleteFile, [true, true, false]),
353 (O::DownloadFileContent, [true, true, false]),
354 ];
355 for (operation, supported) in rows {
356 for (provider, expected) in [P::OpenAi, P::Claude, P::Gemini].into_iter().zip(supported)
357 {
358 let result =
359 request_target(OperationKey::provider(operation, provider), "model", false);
360 assert_eq!(result.is_ok(), expected, "{operation:?} / {provider:?}");
361 }
362 }
363 }
364
365 #[test]
366 fn content_endpoint_support_matrix_is_explicit() {
367 use ContentGenerationKind as K;
368 for kind in [
369 K::OpenAiResponses,
370 K::OpenAiResponsesWebSocket,
371 K::OpenAiChatCompletions,
372 K::ClaudeMessages,
373 K::GeminiGenerateContent,
374 ] {
375 for (operation, stream) in [
376 (Operation::GenerateContent, false),
377 (Operation::StreamGenerateContent, true),
378 ] {
379 let result = request_target(
380 OperationKey::content_generation(operation, kind),
381 "model",
382 stream,
383 );
384 assert_eq!(
385 result.is_ok(),
386 kind != K::OpenAiResponsesWebSocket || stream,
387 "{operation:?} / {kind:?}"
388 );
389 }
390 }
391 }
392
393 #[test]
394 fn rejects_inconsistent_keys_and_stream_flags() {
395 let inconsistent = OperationKey::new_unchecked(
396 Operation::GenerateContent,
397 OperationKind::Provider(Provider::OpenAi),
398 );
399 assert!(matches!(
400 request_target(inconsistent, "model", false),
401 Err(EndpointError::InconsistentOperationKey(_))
402 ));
403
404 let key = OperationKey::content_generation(
405 Operation::StreamGenerateContent,
406 ContentGenerationKind::OpenAiChatCompletions,
407 );
408 assert!(matches!(
409 request_target(key, "model", false),
410 Err(EndpointError::StreamMismatch { .. })
411 ));
412 }
413
414 #[test]
415 fn model_is_a_raw_encoded_component() {
416 let key = OperationKey::provider(Operation::GetModel, Provider::Gemini);
417 let target = request_target(key, "org/model ?", false).unwrap();
418 assert_eq!(target.path, "/v1beta/models/org%2Fmodel%20%3F");
419
420 let key = OperationKey::provider(Operation::ConnectRealtime, Provider::OpenAi);
421 let target = request_target(key, "gpt/a b", false).unwrap();
422 assert_eq!(target.query.as_deref(), Some("model=gpt%2Fa%20b"));
423 }
424}