Skip to main content

gproxy_protocol/protocol/
endpoint.rs

1//! Target endpoint synthesis (M2): the provider-relative method/path/query a
2//! transformed request must hit for a given operation key. Passthrough keeps
3//! the inbound target and never calls this.
4
5use crate::protocol::operation::{
6    ContentGenerationKind, HttpMethod, Operation, OperationKey, OperationKind, Provider,
7};
8
9/// Provider-relative request target for a wired operation.
10#[derive(Debug, Clone, PartialEq, Eq, gproxy_protocol_macros::WireBuilder)]
11#[non_exhaustive]
12pub struct RequestTarget {
13    pub method: HttpMethod,
14    pub path: String,
15    /// Extra query the wire format requires (e.g. gemini `alt=sse`).
16    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
85/// Build the upstream request target for any wired operation key. `model` is
86/// the upstream model id (path-templated providers embed it); `stream` selects
87/// the streaming variant where the wire format distinguishes it by endpoint.
88pub 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        // single-embed form; batch (`:batchEmbedContents`) is a separate op
138        (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        // Imagen-native generation; `gemini-*-image` models route through
147        // generate-content instead.
148        (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        // Video: resource identifiers (video id, operation name, file id)
160        // travel in the `model` argument slot, like `GetModel`.
161        (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        // The poll target is the operation resource name
200        // (`models/<model>/operations/<id>`); slashes are structural.
201        (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
250/// Content-generation targets (POST; gemini selects the verb by `stream`).
251fn 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
298/// Percent-encode one raw path/query component. Model ids passed to this
299/// module are always raw ids, never provider paths or pre-encoded fragments.
300fn 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}