Skip to main content

gproxy_transform/transform/dispatch/
mod.rs

1//! Bytes-level dispatch from a resolved [`TransformPair`] to its typed pair
2//! functions. [`content`] holds the 12 content-generation pairs (M2);
3//! [`other`] holds count_tokens/models/embeddings/images/compact (M2.5).
4//! Streaming is wired for content pairs only.
5
6mod content;
7mod other;
8
9use serde::Serialize;
10use serde::de::DeserializeOwned;
11use serde_json::Value;
12
13use super::{TransformContext, TransformError, TransformPair};
14
15/// Whether the bytes dispatch has arms for this pair.
16pub fn is_wired(pair: TransformPair) -> bool {
17    content::is_content(pair) || other::is_wired(pair)
18}
19
20/// Convert a request body (inbound wire JSON → upstream wire JSON).
21pub fn request_bytes(
22    pair: TransformPair,
23    ctx: &TransformContext,
24    body: &[u8],
25) -> Result<Vec<u8>, TransformError> {
26    if content::is_content(pair) {
27        content::request_bytes(pair, ctx, body)
28    } else {
29        other::request_bytes(pair, ctx, body)
30    }
31}
32
33/// Convert a response body (upstream wire JSON → inbound wire JSON). The pair
34/// here is the REVERSE pair (`resolve(upstream_key, inbound_key)`).
35pub fn response_bytes(
36    pair: TransformPair,
37    ctx: &TransformContext,
38    body: &[u8],
39) -> Result<Vec<u8>, TransformError> {
40    if content::is_content(pair) {
41        content::response_bytes(pair, ctx, body)
42    } else {
43        other::response_bytes(pair, ctx, body)
44    }
45}
46
47/// Convert one decoded stream event (upstream wire JSON value → inbound wire
48/// JSON value). Same reverse-pair convention as [`response_bytes`]. Only
49/// content-generation pairs stream; the other groups are buffered.
50pub fn stream_event_value(
51    pair: TransformPair,
52    ctx: &TransformContext,
53    event: Value,
54) -> Result<Value, TransformError> {
55    if content::is_content(pair) {
56        content::stream_event_value(pair, ctx, event)
57    } else {
58        Err(not_wired(pair))
59    }
60}
61
62fn run<S, T>(
63    f: impl Fn(S, &TransformContext) -> Result<T, TransformError>,
64    ctx: &TransformContext,
65    body: &[u8],
66) -> Result<Vec<u8>, TransformError>
67where
68    S: DeserializeOwned,
69    T: Serialize,
70{
71    let input: S = serde_json::from_slice(body).map_err(|e| TransformError::InvalidInput {
72        reason: format!("decode source body: {e}"),
73    })?;
74    let out = f(input, ctx)?;
75    serde_json::to_vec(&out).map_err(|e| TransformError::Serialization {
76        reason: e.to_string(),
77    })
78}
79
80/// [`run`] for infallible pair functions (plain return, no `Result`).
81fn run_ok<S, T>(
82    f: impl Fn(S, &TransformContext) -> T,
83    ctx: &TransformContext,
84    body: &[u8],
85) -> Result<Vec<u8>, TransformError>
86where
87    S: DeserializeOwned,
88    T: Serialize,
89{
90    run(
91        |input, ctx| Ok::<_, TransformError>(f(input, ctx)),
92        ctx,
93        body,
94    )
95}
96
97fn run_value<S, T>(
98    f: impl Fn(S, &TransformContext) -> Result<T, TransformError>,
99    ctx: &TransformContext,
100    event: Value,
101) -> Result<Value, TransformError>
102where
103    S: DeserializeOwned,
104    T: Serialize,
105{
106    let input: S = serde_json::from_value(event).map_err(|e| TransformError::InvalidInput {
107        reason: format!("decode stream event: {e}"),
108    })?;
109    let out = f(input, ctx)?;
110    serde_json::to_value(&out).map_err(|e| TransformError::Serialization {
111        reason: e.to_string(),
112    })
113}
114
115fn not_wired(pair: TransformPair) -> TransformError {
116    TransformError::InvalidInput {
117        reason: format!("bytes dispatch not wired for {pair:?}"),
118    }
119}
120
121#[cfg(test)]
122mod tests {
123    use super::*;
124    use crate::protocol::{ContentGenerationKind, Operation, OperationKey, Provider};
125
126    #[test]
127    fn claude_to_openai_chat_request_roundtrip() {
128        let source = OperationKey::content_generation(
129            Operation::GenerateContent,
130            ContentGenerationKind::ClaudeMessages,
131        );
132        let target = OperationKey::content_generation(
133            Operation::GenerateContent,
134            ContentGenerationKind::OpenAiChatCompletions,
135        );
136        let ctx = TransformContext::new(source, target);
137        let body = br#"{"model":"m","max_tokens":16,"messages":[{"role":"user","content":"hi"}]}"#;
138        let out = request_bytes(TransformPair::ClaudeMessagesToOpenAiChat, &ctx, body).unwrap();
139        let v: Value = serde_json::from_slice(&out).unwrap();
140        assert_eq!(v["messages"][0]["role"], "user");
141        assert!(v.get("max_tokens").is_some() || v.get("max_completion_tokens").is_some());
142    }
143
144    #[test]
145    fn openai_responses_to_websocket_request_roundtrip() {
146        let source = OperationKey::content_generation(
147            Operation::GenerateContent,
148            ContentGenerationKind::OpenAiResponses,
149        );
150        let target = OperationKey::content_generation(
151            Operation::GenerateContent,
152            ContentGenerationKind::OpenAiResponsesWebSocket,
153        );
154        let pair = crate::transform::resolve(source, target).unwrap();
155        assert!(is_wired(pair));
156
157        let ctx = TransformContext::new(source, target);
158        let body = br#"{"model":"m","input":"hi","stream":true}"#;
159        let out = request_bytes(pair, &ctx, body).unwrap();
160        let v: Value = serde_json::from_slice(&out).unwrap();
161
162        assert_eq!(v["type"], "response.create");
163        assert_eq!(v["model"], "m");
164        assert_eq!(v["stream"], true);
165    }
166
167    #[test]
168    fn claude_to_openai_responses_websocket_request_roundtrip() {
169        let source = OperationKey::content_generation(
170            Operation::GenerateContent,
171            ContentGenerationKind::ClaudeMessages,
172        );
173        let target = OperationKey::content_generation(
174            Operation::GenerateContent,
175            ContentGenerationKind::OpenAiResponsesWebSocket,
176        );
177        let pair = crate::transform::resolve(source, target).unwrap();
178        assert!(is_wired(pair));
179
180        let ctx = TransformContext::new(source, target);
181        let body = br#"{"model":"m","max_tokens":16,"messages":[{"role":"user","content":"hi"}]}"#;
182        let out = request_bytes(pair, &ctx, body).unwrap();
183        let v: Value = serde_json::from_slice(&out).unwrap();
184
185        assert_eq!(v["type"], "response.create");
186        assert_eq!(v["model"], "m");
187        assert!(v.get("input").is_some());
188    }
189
190    #[test]
191    fn claude_to_openai_count_tokens_request_roundtrip() {
192        let source = OperationKey::provider(Operation::CountTokens, Provider::Claude);
193        let target = OperationKey::provider(Operation::CountTokens, Provider::OpenAi);
194        let ctx = TransformContext::new(source, target);
195        let body = br#"{"model":"m","messages":[{"role":"user","content":"hi"}]}"#;
196        let out = request_bytes(TransformPair::ClaudeToOpenAiCountTokens, &ctx, body).unwrap();
197        let v: Value = serde_json::from_slice(&out).unwrap();
198        assert_eq!(v["model"], "m");
199        assert!(v.get("input").is_some());
200    }
201
202    #[test]
203    fn compact_to_responses_is_resolved_and_wired() {
204        let source = OperationKey::provider(Operation::CompactContent, Provider::OpenAi);
205        let target = OperationKey::content_generation(
206            Operation::GenerateContent,
207            ContentGenerationKind::OpenAiResponses,
208        );
209        let pair = crate::transform::resolve(source, target).unwrap();
210        assert_eq!(pair, TransformPair::OpenAiCompactToOpenAiResponses);
211        assert!(is_wired(pair));
212
213        let ctx = TransformContext::new(source, target);
214        let body = br#"{"model":"m","input":"summarize this"}"#;
215        let out = request_bytes(pair, &ctx, body).unwrap();
216        let value: Value = serde_json::from_slice(&out).unwrap();
217        assert_eq!(value["model"], "m");
218        assert!(value.get("input").is_some());
219    }
220
221    #[test]
222    fn openai_to_claude_models_list_response_roundtrip() {
223        let source = OperationKey::provider(Operation::ListModels, Provider::OpenAi);
224        let target = OperationKey::provider(Operation::ListModels, Provider::Claude);
225        let ctx = TransformContext::new(source, target);
226        let body = br#"{"object":"list","data":[{"id":"gpt-x","created":1,"object":"model","owned_by":"openai"}]}"#;
227        let out = response_bytes(TransformPair::OpenAiToClaudeModels, &ctx, body).unwrap();
228        let v: Value = serde_json::from_slice(&out).unwrap();
229        assert_eq!(v["data"][0]["id"], "gpt-x");
230    }
231}