1use agent_base::ToolMetadata as AgentToolMetadata;
28use serde::{Deserialize, Serialize};
29use serde_json::Value;
30
31pub const PROTOCOL_VERSION: u32 = 1;
32
33#[derive(Debug, Clone, Serialize, Deserialize)]
39pub struct ToolMetadata {
40 pub name: String,
41 pub description: String,
42 pub origin: String,
43 pub version: String,
44 pub requirements: Vec<String>,
45}
46
47impl From<AgentToolMetadata> for ToolMetadata {
48 fn from(m: AgentToolMetadata) -> Self {
49 Self {
50 name: m.name,
51 description: m.description,
52 origin: m.origin,
53 version: m.version,
54 requirements: m.requirements,
55 }
56 }
57}
58
59#[derive(Debug, Deserialize)]
62#[serde(tag = "type", rename_all = "snake_case")]
63pub enum IncomingMessage {
64 RegisterTool {
65 name: String,
66 description: String,
67 parameters: Value,
68 },
69 CreateSession {
70 #[serde(default)]
71 session_id: Option<String>,
72 },
73 Run {
74 #[serde(default)]
75 session_id: String,
76 query: String,
77 #[serde(default)]
78 config: Option<RunConfig>,
79 },
80 ToolResult {
81 call_id: String,
82 summary: String,
83 #[serde(default)]
84 raw: Option<Value>,
85 #[serde(default)]
86 control_flow: Option<String>,
87 },
88 Cancel {
89 #[serde(default)]
90 session_id: String,
91 },
92 ListTools {},
93}
94
95#[derive(Debug, Deserialize, Default)]
96pub struct RunConfig {
97 #[serde(default)]
98 pub model: Option<String>,
99 #[serde(default)]
100 pub api_key: Option<String>,
101 #[serde(default)]
102 pub base_url: Option<String>,
103 pub enable_thinking: Option<bool>,
104 pub thinking_budget: Option<u64>,
105 pub thinking_effort: Option<String>,
106 pub max_tool_calls_per_turn: Option<usize>,
107 pub max_consecutive_failures: Option<usize>,
108 pub max_turns: Option<u32>,
109}
110
111#[derive(Debug, Serialize)]
114#[serde(tag = "type", rename_all = "snake_case")]
115pub enum OutgoingMessage {
116 Hello {
117 protocol_version: u32,
118 server_name: String,
119 server_version: String,
120 },
121 SessionCreated {
122 session_id: Option<String>,
123 internal_id: u64,
124 },
125 Event {
126 seq: u64,
127 #[serde(flatten)]
128 event: Value,
129 },
130 ToolCall {
131 seq: u64,
132 call_id: String,
133 name: String,
134 args: Value,
135 },
136 ToolRegistered {
137 name: String,
138 ok: bool,
139 },
140 Done {
141 seq: u64,
142 outcome: String,
143 #[serde(skip_serializing_if = "Option::is_none")]
144 error: Option<String>,
145 #[serde(skip_serializing_if = "Option::is_none")]
146 turns: Option<u32>,
147 },
148 Error {
149 code: String,
150 message: String,
151 #[serde(skip_serializing_if = "Option::is_none")]
152 detail: Option<Value>,
153 },
154 ToolsListed {
155 tools: Vec<ToolMetadata>,
156 },
157}
158
159#[cfg(test)]
160mod tests {
161 use super::*;
162 use agent_base::ToolMetadata as AgentToolMetadata;
163
164 #[test]
167 fn test_deserialize_all_incoming_variants() {
168 let json = r#"{"type":"register_tool","name":"shell","description":"run shell","parameters":{}}"#;
170 let msg: IncomingMessage = serde_json::from_str(json).unwrap();
171 assert!(matches!(msg, IncomingMessage::RegisterTool { .. }));
172
173 let json = r#"{"type":"create_session","session_id":"ext-1"}"#;
175 let msg: IncomingMessage = serde_json::from_str(json).unwrap();
176 assert!(matches!(msg, IncomingMessage::CreateSession { .. }));
177
178 let json = r#"{"type":"create_session"}"#;
180 let msg: IncomingMessage = serde_json::from_str(json).unwrap();
181 assert!(matches!(msg, IncomingMessage::CreateSession { session_id: None }));
182
183 let json = r#"{"type":"run","session_id":"abc","query":"hello"}"#;
185 let msg: IncomingMessage = serde_json::from_str(json).unwrap();
186 assert!(matches!(msg, IncomingMessage::Run { .. }));
187
188 let json = r#"{"type":"tool_result","call_id":"c1","summary":"done"}"#;
190 let msg: IncomingMessage = serde_json::from_str(json).unwrap();
191 assert!(matches!(msg, IncomingMessage::ToolResult { .. }));
192
193 let json = r#"{"type":"cancel","session_id":"abc"}"#;
195 let msg: IncomingMessage = serde_json::from_str(json).unwrap();
196 assert!(matches!(msg, IncomingMessage::Cancel { .. }));
197
198 let json = r#"{"type":"list_tools"}"#;
200 let msg: IncomingMessage = serde_json::from_str(json).unwrap();
201 assert!(matches!(msg, IncomingMessage::ListTools {}));
202 }
203
204 #[test]
205 fn test_run_config_all_fields_populated() {
206 let json = r#"{
207 "type":"run",
208 "session_id":"s1",
209 "query":"test",
210 "config":{
211 "model":"gpt-4",
212 "api_key":"sk-xxx",
213 "base_url":"https://example.com/v1",
214 "enable_thinking":true,
215 "thinking_budget":32000,
216 "thinking_effort":"high",
217 "max_tool_calls_per_turn":10,
218 "max_consecutive_failures":3,
219 "max_turns":5
220 }
221 }"#;
222 let msg: IncomingMessage = serde_json::from_str(json).unwrap();
223 if let IncomingMessage::Run { config: Some(cfg), .. } = &msg {
224 assert_eq!(cfg.model.as_deref(), Some("gpt-4"));
225 assert_eq!(cfg.api_key.as_deref(), Some("sk-xxx"));
226 assert_eq!(cfg.base_url.as_deref(), Some("https://example.com/v1"));
227 assert_eq!(cfg.enable_thinking, Some(true));
228 assert_eq!(cfg.thinking_budget, Some(32000));
229 assert_eq!(cfg.thinking_effort.as_deref(), Some("high"));
230 assert_eq!(cfg.max_tool_calls_per_turn, Some(10));
231 assert_eq!(cfg.max_consecutive_failures, Some(3));
232 assert_eq!(cfg.max_turns, Some(5));
233 } else {
234 panic!("expected Run with config");
235 }
236 }
237
238 #[test]
239 fn test_run_config_default_all_none() {
240 let cfg: RunConfig = serde_json::from_str("{}").unwrap();
241 assert!(cfg.model.is_none());
242 assert!(cfg.api_key.is_none());
243 assert!(cfg.base_url.is_none());
244 assert!(cfg.enable_thinking.is_none());
245 assert!(cfg.thinking_budget.is_none());
246 assert!(cfg.thinking_effort.is_none());
247 assert!(cfg.max_tool_calls_per_turn.is_none());
248 assert!(cfg.max_consecutive_failures.is_none());
249 assert!(cfg.max_turns.is_none());
250 }
251
252 #[test]
253 fn test_run_without_config_uses_none() {
254 let json = r#"{"type":"run","session_id":"s1","query":"test"}"#;
255 let msg: IncomingMessage = serde_json::from_str(json).unwrap();
256 if let IncomingMessage::Run { config, .. } = &msg {
257 assert!(config.is_none());
258 } else {
259 panic!("expected Run");
260 }
261 }
262
263 #[test]
264 fn test_unknown_fields_ignored() {
265 let json = r#"{"type":"list_tools","extra_field":"should-be-ignored","nested":{"a":1}}"#;
266 let msg: IncomingMessage = serde_json::from_str(json).unwrap();
267 assert!(matches!(msg, IncomingMessage::ListTools {}));
268 }
269
270 #[test]
271 fn test_missing_type_field_errors() {
272 let json = r#"{"session_id":"abc","query":"hello"}"#;
273 let result = serde_json::from_str::<IncomingMessage>(json);
274 assert!(result.is_err());
275 }
276
277 #[test]
278 fn test_deserialize_tool_result_with_all_fields() {
279 let json = r#"{"type":"tool_result","call_id":"c1","summary":"done","raw":{"output":"hello"},"control_flow":"break"}"#;
280 let msg: IncomingMessage = serde_json::from_str(json).unwrap();
281 if let IncomingMessage::ToolResult { call_id, summary, raw, control_flow } = &msg {
282 assert_eq!(call_id, "c1");
283 assert_eq!(summary, "done");
284 assert_eq!(raw.as_ref().and_then(|v| v.get("output")).and_then(|v| v.as_str()), Some("hello"));
285 assert_eq!(control_flow.as_deref(), Some("break"));
286 } else {
287 panic!("expected ToolResult");
288 }
289 }
290
291 #[test]
294 fn test_serialize_all_outgoing_variants() {
295 let msg = OutgoingMessage::Hello {
297 protocol_version: 1,
298 server_name: "phi".into(),
299 server_version: "0.2.6".into(),
300 };
301 let json = serde_json::to_value(&msg).unwrap();
302 assert_eq!(json["type"], "hello");
303 assert_eq!(json["protocol_version"], 1);
304
305 let msg = OutgoingMessage::SessionCreated { session_id: Some("ext-1".into()), internal_id: 42 };
307 let json = serde_json::to_value(&msg).unwrap();
308 assert_eq!(json["type"], "session_created");
309 assert_eq!(json["internal_id"], 42);
310
311 let msg = OutgoingMessage::Event {
313 seq: 1,
314 event: serde_json::json!({"type":"text_delta","text":"hi"}),
315 };
316 let json = serde_json::to_value(&msg).unwrap();
317 assert_eq!(json["seq"], 1);
318 assert_eq!(json["text"], "hi"); let msg = OutgoingMessage::ToolCall {
322 seq: 2,
323 call_id: "c1".into(),
324 name: "shell".into(),
325 args: serde_json::json!({"cmd":"ls"}),
326 };
327 let json = serde_json::to_value(&msg).unwrap();
328 assert_eq!(json["type"], "tool_call");
329 assert_eq!(json["call_id"], "c1");
330
331 let msg = OutgoingMessage::ToolRegistered { name: "shell".into(), ok: true };
333 let json = serde_json::to_value(&msg).unwrap();
334 assert_eq!(json["type"], "tool_registered");
335 assert!(json["ok"].as_bool().unwrap());
336
337 let msg = OutgoingMessage::Done { seq: 3, outcome: "completed".into(), error: None, turns: Some(1) };
339 let json = serde_json::to_value(&msg).unwrap();
340 assert_eq!(json["type"], "done");
341 assert_eq!(json["outcome"], "completed");
342 assert!(json.get("error").is_none());
343
344 let msg = OutgoingMessage::Done {
346 seq: 4,
347 outcome: "failed".into(),
348 error: Some("something went wrong".into()),
349 turns: None,
350 };
351 let json = serde_json::to_value(&msg).unwrap();
352 assert_eq!(json["error"], "something went wrong");
353 assert!(json.get("turns").is_none());
354
355 let msg = OutgoingMessage::Error {
357 code: "E001".into(),
358 message: "bad request".into(),
359 detail: None,
360 };
361 let json = serde_json::to_value(&msg).unwrap();
362 assert_eq!(json["type"], "error");
363 assert!(json.get("detail").is_none());
364
365 let msg = OutgoingMessage::ToolsListed { tools: vec![] };
367 let json = serde_json::to_value(&msg).unwrap();
368 assert_eq!(json["type"], "tools_listed");
369 }
370
371 #[test]
374 fn test_tool_metadata_from_agent_tool_metadata() {
375 let am = AgentToolMetadata {
376 name: "shell".into(),
377 description: "Run shell commands".into(),
378 origin: "phi-tools".into(),
379 version: "1.0.0".into(),
380 requirements: vec!["bash".into()],
381 };
382 let tm = ToolMetadata::from(am);
383 assert_eq!(tm.name, "shell");
384 assert_eq!(tm.description, "Run shell commands");
385 assert_eq!(tm.origin, "phi-tools");
386 assert_eq!(tm.version, "1.0.0");
387 assert_eq!(tm.requirements, vec!["bash"]);
388 }
389
390 #[test]
391 fn test_tool_metadata_round_trip() {
392 let tm = ToolMetadata {
393 name: "shell".into(),
394 description: "desc".into(),
395 origin: "phi".into(),
396 version: "1.0".into(),
397 requirements: vec!["bash".into(), "zsh".into()],
398 };
399 let json = serde_json::to_string(&tm).unwrap();
400 let back: ToolMetadata = serde_json::from_str(&json).unwrap();
401 assert_eq!(back.name, tm.name);
402 assert_eq!(back.description, tm.description);
403 assert_eq!(back.origin, tm.origin);
404 assert_eq!(back.version, tm.version);
405 assert_eq!(back.requirements, tm.requirements);
406 }
407}