Skip to main content

gproxy_transform/transform/dispatch/
mod.rs

1//! Bytes-level dispatch from a resolved [`TransformPair`] to its typed pair
2//! functions. The private `content` module holds content-generation pairs;
3//! `other` holds count_tokens/models/embeddings/images/compact.
4//! Streaming is wired for content pairs only.
5
6mod content;
7mod other;
8
9use serde::Serialize;
10use serde::de::DeserializeOwned;
11
12use super::{TransformContext, TransformError, TransformOutput, TransformPair};
13
14/// Whether the bytes dispatch has arms for this pair.
15pub fn is_wired(pair: TransformPair) -> bool {
16    content::is_content(pair) || other::is_wired(pair)
17}
18
19/// Convert a request body (inbound wire JSON → upstream wire JSON).
20pub fn request_bytes(
21    pair: TransformPair,
22    ctx: &TransformContext,
23    body: &[u8],
24) -> Result<Vec<u8>, TransformError> {
25    validate_pair(pair, ctx)?;
26    ctx.scope(|| {
27        if content::is_content(pair) {
28            content::request_bytes(pair, ctx, body)
29        } else {
30            other::request_bytes(pair, ctx, body)
31        }
32    })
33}
34
35/// Convert a request and return non-fatal semantic-loss diagnostics.
36pub fn request_bytes_detailed(
37    pair: TransformPair,
38    ctx: &TransformContext,
39    body: &[u8],
40) -> Result<TransformOutput<Vec<u8>>, TransformError> {
41    let scoped = ctx.isolated();
42    let value = request_bytes(pair, &scoped, body)?;
43    Ok(TransformOutput::new(value, scoped.take_diagnostics()))
44}
45
46/// Convert a response body (upstream wire JSON → inbound wire JSON). The pair
47/// here is the REVERSE pair (`resolve(upstream_key, inbound_key)`).
48pub fn response_bytes(
49    pair: TransformPair,
50    ctx: &TransformContext,
51    body: &[u8],
52) -> Result<Vec<u8>, TransformError> {
53    validate_pair(pair, ctx)?;
54    ctx.scope(|| {
55        if content::is_content(pair) {
56            content::response_bytes(pair, ctx, body)
57        } else {
58            other::response_bytes(pair, ctx, body)
59        }
60    })
61}
62
63/// Convert a response and return non-fatal semantic-loss diagnostics.
64pub fn response_bytes_detailed(
65    pair: TransformPair,
66    ctx: &TransformContext,
67    body: &[u8],
68) -> Result<TransformOutput<Vec<u8>>, TransformError> {
69    let scoped = ctx.isolated();
70    let value = response_bytes(pair, &scoped, body)?;
71    Ok(TransformOutput::new(value, scoped.take_diagnostics()))
72}
73
74/// One converted stream event: pre-encoded inbound frame payload, or the
75/// typed Responses event when the inbound side runs the aggregation state
76/// machine.
77pub enum StreamEventOut {
78    Encoded { event: Option<String>, data: String },
79    Responses(Box<crate::protocol::openai::ResponseStreamEvent>),
80}
81
82pub(crate) struct StreamEventBatch {
83    pub events: Vec<StreamEventOut>,
84    pub terminal: bool,
85}
86
87/// Stateful `0..N` stream-event converter for one resolved pair.
88///
89/// Create one converter per upstream response and retain it until
90/// [`finish`](Self::finish). This preserves pair-specific state such as tool
91/// call arguments split across multiple frames.
92pub struct StreamConverter {
93    inner: content::ContentStreamConverter,
94}
95
96impl StreamConverter {
97    pub fn new(pair: TransformPair, ctx: TransformContext) -> Result<Self, TransformError> {
98        validate_pair(pair, &ctx)?;
99        Ok(Self {
100            inner: content::ContentStreamConverter::new(pair, ctx)?,
101        })
102    }
103
104    /// Convert one decoded upstream event into zero or more inbound events.
105    pub fn push(&mut self, data: &str) -> Result<Vec<StreamEventOut>, TransformError> {
106        Ok(self.push_detailed(data)?.value)
107    }
108
109    /// Convert one event and return semantic-loss diagnostics produced by it.
110    pub fn push_detailed(
111        &mut self,
112        data: &str,
113    ) -> Result<TransformOutput<Vec<StreamEventOut>>, TransformError> {
114        let output = self.push_detailed_with_status(data)?;
115        Ok(TransformOutput::new(
116            output.value.events,
117            output.diagnostics,
118        ))
119    }
120
121    pub(crate) fn push_detailed_with_status(
122        &mut self,
123        data: &str,
124    ) -> Result<TransformOutput<StreamEventBatch>, TransformError> {
125        let output = self.inner.push(data)?;
126        Ok(TransformOutput::new(
127            StreamEventBatch {
128                events: output.events,
129                terminal: output.terminal,
130            },
131            self.inner.take_diagnostics(),
132        ))
133    }
134
135    /// Flush pair-specific state into zero or more final inbound events.
136    pub fn finish(&mut self) -> Result<Vec<StreamEventOut>, TransformError> {
137        Ok(self.finish_detailed()?.value)
138    }
139
140    /// Flush state and return any final semantic-loss diagnostics.
141    pub fn finish_detailed(
142        &mut self,
143    ) -> Result<TransformOutput<Vec<StreamEventOut>>, TransformError> {
144        let value = self.inner.finish()?;
145        Ok(TransformOutput::new(value, self.inner.take_diagnostics()))
146    }
147}
148
149/// Convert one decoded stream event (upstream wire JSON text → inbound event).
150/// Same reverse-pair convention as [`response_bytes`]. Only content-generation
151/// pairs stream; the other groups are buffered. This convenience call retains
152/// no cross-frame state; use [`StreamConverter`] for an actual response stream.
153pub fn stream_event(
154    pair: TransformPair,
155    ctx: &TransformContext,
156    data: &str,
157) -> Result<Vec<StreamEventOut>, TransformError> {
158    if content::is_content(pair) {
159        let mut converter = StreamConverter::new(pair, ctx.clone())?;
160        converter.push(data)
161    } else {
162        Err(not_wired(pair))
163    }
164}
165
166fn validate_pair(pair: TransformPair, ctx: &TransformContext) -> Result<(), TransformError> {
167    let resolved = super::resolve(ctx.source, ctx.target)?;
168    if resolved == pair {
169        Ok(())
170    } else {
171        Err(TransformError::InvalidInput {
172            reason: format!(
173                "transform pair {pair:?} does not match context {:?} -> {:?} (resolved {resolved:?})",
174                ctx.source, ctx.target
175            ),
176        })
177    }
178}
179
180fn run<S, T>(
181    f: impl Fn(S, &TransformContext) -> Result<T, TransformError>,
182    ctx: &TransformContext,
183    body: &[u8],
184) -> Result<Vec<u8>, TransformError>
185where
186    S: DeserializeOwned,
187    T: Serialize,
188{
189    let input: S = serde_json::from_slice(body).map_err(|e| TransformError::InvalidInput {
190        reason: format!("decode source body: {e}"),
191    })?;
192    let out = f(input, ctx)?;
193    serde_json::to_vec(&out).map_err(|e| TransformError::Serialization {
194        reason: e.to_string(),
195    })
196}
197
198/// [`run`] for infallible pair functions (plain return, no `Result`).
199fn run_ok<S, T>(
200    f: impl Fn(S, &TransformContext) -> T,
201    ctx: &TransformContext,
202    body: &[u8],
203) -> Result<Vec<u8>, TransformError>
204where
205    S: DeserializeOwned,
206    T: Serialize,
207{
208    run(
209        |input, ctx| Ok::<_, TransformError>(f(input, ctx)),
210        ctx,
211        body,
212    )
213}
214
215fn not_wired(pair: TransformPair) -> TransformError {
216    TransformError::InvalidInput {
217        reason: format!("bytes dispatch not wired for {pair:?}"),
218    }
219}
220
221#[cfg(test)]
222mod tests {
223    use serde_json::Value;
224
225    use super::*;
226    use crate::protocol::{ContentGenerationKind, Operation, OperationKey, Provider};
227    use crate::transform::TransformDiagnosticKind;
228
229    #[test]
230    fn claude_to_openai_chat_request_roundtrip() {
231        let source = OperationKey::content_generation(
232            Operation::GenerateContent,
233            ContentGenerationKind::ClaudeMessages,
234        );
235        let target = OperationKey::content_generation(
236            Operation::GenerateContent,
237            ContentGenerationKind::OpenAiChatCompletions,
238        );
239        let ctx = TransformContext::new(source, target);
240        let body = br#"{"model":"m","max_tokens":16,"messages":[{"role":"user","content":"hi"}]}"#;
241        let out = request_bytes(TransformPair::ClaudeMessagesToOpenAiChat, &ctx, body).unwrap();
242        let v: Value = serde_json::from_slice(&out).unwrap();
243        assert_eq!(v["messages"][0]["role"], "user");
244        assert!(v.get("max_tokens").is_some() || v.get("max_completion_tokens").is_some());
245    }
246
247    #[test]
248    fn detailed_request_returns_structured_semantic_loss() {
249        let source = OperationKey::content_generation(
250            Operation::GenerateContent,
251            ContentGenerationKind::OpenAiChatCompletions,
252        );
253        let target = OperationKey::content_generation(
254            Operation::GenerateContent,
255            ContentGenerationKind::ClaudeMessages,
256        );
257        let pair = crate::transform::resolve(source, target).unwrap();
258        let ctx = TransformContext::new(source, target);
259        let body = br#"{
260            "model":"m",
261            "messages":[{"role":"user","content":[{
262                "type":"text",
263                "text":"",
264                "prompt_cache_breakpoint":{"mode":"explicit"}
265            }]}]
266        }"#;
267
268        let output = request_bytes_detailed(pair, &ctx, body).unwrap();
269
270        assert_eq!(output.diagnostics.len(), 1);
271        assert_eq!(
272            output.diagnostics[0].kind,
273            TransformDiagnosticKind::LossyField
274        );
275        assert_eq!(
276            output.diagnostics[0].field,
277            "messages[].content[].text.prompt_cache_breakpoint"
278        );
279        assert!(ctx.diagnostics().is_empty(), "detailed calls are isolated");
280    }
281
282    #[test]
283    fn openai_responses_to_websocket_request_roundtrip() {
284        let source = OperationKey::content_generation(
285            Operation::GenerateContent,
286            ContentGenerationKind::OpenAiResponses,
287        );
288        let target = OperationKey::content_generation(
289            Operation::GenerateContent,
290            ContentGenerationKind::OpenAiResponsesWebSocket,
291        );
292        let pair = crate::transform::resolve(source, target).unwrap();
293        assert!(is_wired(pair));
294
295        let ctx = TransformContext::new(source, target);
296        let body = br#"{"model":"m","input":"hi","stream":true}"#;
297        let out = request_bytes(pair, &ctx, body).unwrap();
298        let v: Value = serde_json::from_slice(&out).unwrap();
299
300        assert_eq!(v["type"], "response.create");
301        assert_eq!(v["model"], "m");
302        assert_eq!(v["stream"], true);
303    }
304
305    #[test]
306    fn claude_to_openai_responses_websocket_request_roundtrip() {
307        let source = OperationKey::content_generation(
308            Operation::GenerateContent,
309            ContentGenerationKind::ClaudeMessages,
310        );
311        let target = OperationKey::content_generation(
312            Operation::GenerateContent,
313            ContentGenerationKind::OpenAiResponsesWebSocket,
314        );
315        let pair = crate::transform::resolve(source, target).unwrap();
316        assert!(is_wired(pair));
317
318        let ctx = TransformContext::new(source, target);
319        let body = br#"{"model":"m","max_tokens":16,"messages":[{"role":"user","content":"hi"}]}"#;
320        let out = request_bytes(pair, &ctx, body).unwrap();
321        let v: Value = serde_json::from_slice(&out).unwrap();
322
323        assert_eq!(v["type"], "response.create");
324        assert_eq!(v["model"], "m");
325        assert!(v.get("input").is_some());
326    }
327
328    #[test]
329    fn claude_to_openai_count_tokens_request_roundtrip() {
330        let source = OperationKey::provider(Operation::CountTokens, Provider::Claude);
331        let target = OperationKey::provider(Operation::CountTokens, Provider::OpenAi);
332        let ctx = TransformContext::new(source, target);
333        let body = br#"{"model":"m","messages":[{"role":"user","content":"hi"}]}"#;
334        let out = request_bytes(TransformPair::ClaudeToOpenAiCountTokens, &ctx, body).unwrap();
335        let v: Value = serde_json::from_slice(&out).unwrap();
336        assert_eq!(v["model"], "m");
337        assert!(v.get("input").is_some());
338    }
339
340    #[test]
341    fn compact_to_responses_is_resolved_and_wired() {
342        let source = OperationKey::provider(Operation::CompactContent, Provider::OpenAi);
343        let target = OperationKey::content_generation(
344            Operation::GenerateContent,
345            ContentGenerationKind::OpenAiResponses,
346        );
347        let pair = crate::transform::resolve(source, target).unwrap();
348        assert_eq!(pair, TransformPair::OpenAiCompactToOpenAiResponses);
349        assert!(is_wired(pair));
350
351        let ctx = TransformContext::new(source, target);
352        let body = br#"{"model":"m","input":"summarize this"}"#;
353        let out = request_bytes(pair, &ctx, body).unwrap();
354        let value: Value = serde_json::from_slice(&out).unwrap();
355        assert_eq!(value["model"], "m");
356        assert!(value.get("input").is_some());
357    }
358
359    #[test]
360    fn openai_to_claude_models_list_response_roundtrip() {
361        let source = OperationKey::provider(Operation::ListModels, Provider::OpenAi);
362        let target = OperationKey::provider(Operation::ListModels, Provider::Claude);
363        let ctx = TransformContext::new(source, target);
364        let body = br#"{"object":"list","data":[{"id":"gpt-x","created":1,"object":"model","owned_by":"openai"}]}"#;
365        let out = response_bytes(TransformPair::OpenAiToClaudeModels, &ctx, body).unwrap();
366        let v: Value = serde_json::from_slice(&out).unwrap();
367        assert_eq!(v["data"][0]["id"], "gpt-x");
368    }
369}