gproxy_transform/transform/dispatch/
mod.rs1mod content;
7mod other;
8
9use serde::Serialize;
10use serde::de::DeserializeOwned;
11use serde_json::Value;
12
13use super::{TransformContext, TransformError, TransformPair};
14
15pub fn is_wired(pair: TransformPair) -> bool {
17 content::is_content(pair) || other::is_wired(pair)
18}
19
20pub 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
33pub 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
47pub 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
80fn 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}