Skip to main content

shore_protocol/
client_msg.rs

1use serde::{Deserialize, Serialize};
2
3/// Client hello — sent once after connect.
4#[derive(Serialize, Deserialize, Debug, Clone)]
5pub struct ClientHello {
6    pub client_type: String,
7    pub client_name: String,
8    #[serde(default)]
9    pub capabilities: Vec<String>,
10    /// Which character this client wants to talk to.
11    #[serde(default, skip_serializing_if = "Option::is_none")]
12    pub character: Option<String>,
13}
14
15/// One-shot parameter overrides for a single message.
16#[derive(Serialize, Deserialize, Debug, Clone, Default)]
17pub struct MessageOverrides {
18    #[serde(skip_serializing_if = "Option::is_none")]
19    pub temperature: Option<f64>,
20    #[serde(skip_serializing_if = "Option::is_none")]
21    pub top_p: Option<f64>,
22    /// Enable extended thinking with the given budget (in tokens).
23    /// `Some(n)` enables thinking with budget `n`; omitted = use model default.
24    #[serde(skip_serializing_if = "Option::is_none")]
25    pub thinking_budget: Option<u32>,
26}
27
28/// A base64-encoded image uploaded by the client.
29#[derive(Serialize, Deserialize, Debug, Clone)]
30pub struct ImageUpload {
31    pub filename: String,
32    /// Base64-encoded image file bytes.
33    pub data: String,
34}
35
36/// Send a user message.
37#[derive(Serialize, Deserialize, Debug, Clone)]
38pub struct ClientMessageBody {
39    #[serde(skip_serializing_if = "Option::is_none")]
40    pub rid: Option<String>,
41    pub text: String,
42    #[serde(default)]
43    pub stream: bool,
44    /// Legacy: file paths (only works when client and server share a filesystem).
45    #[serde(default)]
46    pub images: Vec<String>,
47    /// Preferred: base64-encoded image data (works across machines).
48    #[serde(default, skip_serializing_if = "Vec::is_empty")]
49    pub image_data: Vec<ImageUpload>,
50    #[serde(skip_serializing_if = "Option::is_none")]
51    pub absence_seconds: Option<u64>,
52    #[serde(default, skip_serializing_if = "Option::is_none")]
53    pub overrides: Option<MessageOverrides>,
54}
55
56/// Regenerate last response.
57#[derive(Serialize, Deserialize, Debug, Clone)]
58pub struct Regen {
59    #[serde(skip_serializing_if = "Option::is_none")]
60    pub rid: Option<String>,
61    #[serde(default)]
62    pub stream: bool,
63    #[serde(skip_serializing_if = "Option::is_none")]
64    pub guidance: Option<String>,
65}
66
67/// Execute a server command.
68#[derive(Serialize, Deserialize, Debug, Clone)]
69pub struct Command {
70    #[serde(skip_serializing_if = "Option::is_none")]
71    pub rid: Option<String>,
72    pub name: String,
73    #[serde(default)]
74    pub args: serde_json::Value,
75}
76
77/// Cancel an in-progress generation.
78#[derive(Serialize, Deserialize, Debug, Clone)]
79pub struct Cancel {}
80
81/// All client → server message types, tagged by "type".
82#[derive(Serialize, Deserialize, Debug, Clone)]
83#[serde(tag = "type", rename_all = "snake_case")]
84pub enum ClientMessage {
85    Hello(ClientHello),
86    Message(ClientMessageBody),
87    Regen(Regen),
88    Command(Command),
89    Cancel(Cancel),
90}
91
92#[cfg(test)]
93mod tests {
94    use super::*;
95
96    #[test]
97    fn cancel_serialization_roundtrip() {
98        let msg = ClientMessage::Cancel(Cancel {});
99        let json = serde_json::to_value(&msg).unwrap();
100        assert_eq!(json["type"], "cancel");
101
102        let roundtrip: ClientMessage = serde_json::from_value(json).unwrap();
103        assert!(matches!(roundtrip, ClientMessage::Cancel(_)));
104    }
105
106    #[test]
107    fn message_overrides_with_values() {
108        let overrides = MessageOverrides {
109            temperature: Some(0.8),
110            top_p: Some(0.95),
111            thinking_budget: Some(4096),
112        };
113        let json = serde_json::to_value(&overrides).unwrap();
114        assert_eq!(json["temperature"], 0.8);
115        assert_eq!(json["top_p"], 0.95);
116        assert_eq!(json["thinking_budget"], 4096);
117    }
118
119    #[test]
120    fn message_overrides_none_fields_omitted() {
121        let overrides = MessageOverrides::default();
122        let json = serde_json::to_value(&overrides).unwrap();
123        assert!(json.get("temperature").is_none());
124        assert!(json.get("top_p").is_none());
125        assert!(json.get("thinking_budget").is_none());
126    }
127
128    #[test]
129    fn message_overrides_partial_fields() {
130        let overrides = MessageOverrides {
131            temperature: Some(0.5),
132            top_p: None,
133            thinking_budget: None,
134        };
135        let json = serde_json::to_value(&overrides).unwrap();
136        assert_eq!(json["temperature"], 0.5);
137        assert!(json.get("top_p").is_none());
138    }
139
140    #[test]
141    fn client_message_body_with_overrides_roundtrip() {
142        let body = ClientMessageBody {
143            rid: Some("r1".into()),
144            text: "hello".into(),
145            stream: true,
146            images: vec![],
147            image_data: vec![],
148            absence_seconds: None,
149            overrides: Some(MessageOverrides {
150                temperature: Some(0.7),
151                top_p: None,
152                thinking_budget: Some(2048),
153            }),
154        };
155        let msg = ClientMessage::Message(body);
156        let json = serde_json::to_value(&msg).unwrap();
157        assert_eq!(json["type"], "message");
158        assert_eq!(json["overrides"]["temperature"], 0.7);
159        assert_eq!(json["overrides"]["thinking_budget"], 2048);
160        assert!(json["overrides"].get("top_p").is_none());
161
162        let roundtrip: ClientMessage = serde_json::from_value(json).unwrap();
163        match roundtrip {
164            ClientMessage::Message(b) => {
165                let o = b.overrides.unwrap();
166                assert_eq!(o.temperature, Some(0.7));
167                assert_eq!(o.thinking_budget, Some(2048));
168                assert_eq!(o.top_p, None);
169            }
170            _ => panic!("wrong variant"),
171        }
172    }
173
174    #[test]
175    fn client_message_body_without_overrides() {
176        let body = ClientMessageBody {
177            rid: None,
178
179            text: "hi".into(),
180            stream: false,
181            images: vec![],
182            image_data: vec![],
183            absence_seconds: None,
184            overrides: None,
185        };
186        let json = serde_json::to_value(&body).unwrap();
187        assert!(json.get("overrides").is_none());
188    }
189}