Skip to main content

gproxy_protocol/protocol/openai/generate_content/
responses_websocket.rs

1use serde::{Deserialize, Serialize};
2
3use super::super::common::*;
4use super::{ResponseCreateRequest, ResponseInput, ResponseStreamEvent};
5
6pub type ResponseWebSocketWireModel =
7    OpenAiWireModel<ResponseWebSocketRequest, ResponseStreamEvent>;
8
9#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
10#[serde(tag = "type")]
11#[allow(clippy::large_enum_variant)]
12#[non_exhaustive]
13pub enum ResponseWebSocketRequest {
14    #[serde(rename = "response.create")]
15    ResponseCreate(ResponseCreateWebSocketRequest),
16    #[serde(rename = "response.steer")]
17    ResponseSteer(ResponseSteerWebSocketRequest),
18}
19
20#[derive(
21    Debug, Clone, PartialEq, Serialize, Deserialize, Default, gproxy_protocol_macros::WireBuilder,
22)]
23#[non_exhaustive]
24pub struct ResponseCreateWebSocketRequest {
25    #[serde(flatten)]
26    pub response: ResponseCreateRequest,
27    #[serde(skip_serializing_if = "Option::is_none")]
28    pub generate: Option<bool>,
29    #[serde(skip_serializing_if = "Option::is_none")]
30    pub client_metadata: Option<Metadata>,
31}
32
33#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
34#[non_exhaustive]
35pub struct ResponseSteerWebSocketRequest {
36    pub previous_response_id: String,
37    pub input: ResponseInput,
38    #[serde(
39        default,
40        flatten,
41        skip_serializing_if = "std::collections::BTreeMap::is_empty"
42    )]
43    pub extra: Extra,
44}
45
46#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
47#[non_exhaustive]
48pub struct ResponseSteerReference {
49    pub id: String,
50    pub previous_response_id: String,
51    #[serde(
52        default,
53        flatten,
54        skip_serializing_if = "std::collections::BTreeMap::is_empty"
55    )]
56    pub extra: Extra,
57}
58
59#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
60#[non_exhaustive]
61pub struct ResponseSteerFailure {
62    #[serde(skip_serializing_if = "Option::is_none")]
63    pub id: Option<String>,
64    pub previous_response_id: String,
65    pub input: ResponseInput,
66    #[serde(
67        default,
68        flatten,
69        skip_serializing_if = "std::collections::BTreeMap::is_empty"
70    )]
71    pub extra: Extra,
72}
73
74#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
75#[non_exhaustive]
76pub struct ResponseSteerError {
77    pub code: String,
78    pub message: String,
79    #[serde(rename = "type", skip_serializing_if = "Option::is_none")]
80    pub type_: Option<String>,
81    #[serde(
82        default,
83        flatten,
84        skip_serializing_if = "std::collections::BTreeMap::is_empty"
85    )]
86    pub extra: Extra,
87}
88
89#[cfg(test)]
90mod tests {
91    use serde_json::json;
92
93    use super::*;
94    use crate::protocol::openai::{OpenAiModelId, ResponseInput};
95
96    #[test]
97    fn response_create_frame_round_trips_websocket_fields() {
98        let value = json!({
99            "type": "response.create",
100            "model": "gpt-x",
101            "input": "hello",
102            "stream": true,
103            "generate": false,
104            "client_metadata": {
105                "x-codex-installation-id": "installation-1"
106            }
107        });
108
109        let parsed: ResponseWebSocketRequest = serde_json::from_value(value).unwrap();
110        let ResponseWebSocketRequest::ResponseCreate(frame) = parsed else {
111            panic!("expected response.create")
112        };
113        assert_eq!(
114            frame.response.model,
115            Some(OpenAiModelId::Unknown("gpt-x".to_owned()))
116        );
117        assert_eq!(
118            frame.response.input,
119            Some(ResponseInput::Text("hello".to_owned()))
120        );
121        assert_eq!(frame.response.stream, Some(true));
122        assert_eq!(frame.generate, Some(false));
123        assert_eq!(
124            frame
125                .client_metadata
126                .as_ref()
127                .and_then(|m| m.get("x-codex-installation-id")),
128            Some(&"installation-1".to_owned())
129        );
130
131        let serialized = serde_json::to_value(ResponseWebSocketRequest::ResponseCreate(frame))
132            .expect("serialize websocket response.create");
133        assert_eq!(serialized["type"], json!("response.create"));
134        assert_eq!(serialized["generate"], json!(false));
135    }
136}