oxicode_sdk/middleware/
bridge.rs1use std::sync::Arc;
4use std::sync::atomic::{AtomicBool, Ordering};
5
6use oxicode_agent::AgentHooks;
7
8use crate::middleware::{MiddlewareContext, MiddlewareData, MiddlewarePhase, MiddlewarePipeline};
9
10pub fn build_hooks(
15 pipeline: Arc<MiddlewarePipeline>,
16 agent_id: String,
17 terminate_flag: Arc<AtomicBool>,
18) -> AgentHooks {
19 let before_tool_call = Some(std::boxed::Box::new({
20 let pipeline = Arc::clone(&pipeline);
21 let agent_id = agent_id.clone();
22 let terminate_flag = Arc::clone(&terminate_flag);
23
24 move |ctx: &oxicode_agent::BeforeToolCallContext| -> oxicode_agent::BeforeToolCallResult {
25 let mw_ctx = MiddlewareContext::new(
26 MiddlewarePhase::BeforeTool,
27 &agent_id,
28 MiddlewareData::BeforeTool {
29 tool_name: ctx.tool_name.clone(),
30 params: ctx.args.clone(),
31 },
32 );
33
34 let rt = tokio::runtime::Handle::current();
36 let result = rt.block_on(pipeline.execute(&mw_ctx));
37
38 match result.action {
39 crate::middleware::MiddlewareAction::Continue => {
40 oxicode_agent::BeforeToolCallResult {
41 block: false,
42 reason: None,
43 }
44 }
45 crate::middleware::MiddlewareAction::Block => oxicode_agent::BeforeToolCallResult {
46 block: true,
47 reason: result.reason,
48 },
49 crate::middleware::MiddlewareAction::Terminate => {
50 terminate_flag.store(true, Ordering::SeqCst);
51 oxicode_agent::BeforeToolCallResult {
52 block: true,
53 reason: result.reason,
54 }
55 }
56 }
57 }
58 })
59 as std::boxed::Box<
60 dyn Fn(&oxicode_agent::BeforeToolCallContext) -> oxicode_agent::BeforeToolCallResult
61 + Send
62 + Sync,
63 >);
64
65 let after_tool_call = Some(std::boxed::Box::new({
66 let pipeline = Arc::clone(&pipeline);
67 let agent_id = agent_id.clone();
68 let terminate_flag = Arc::clone(&terminate_flag);
69
70 move |ctx: &oxicode_agent::AfterToolCallContext| -> oxicode_agent::AfterToolCallResult {
71 let mw_ctx = MiddlewareContext::new(
72 MiddlewarePhase::AfterTool,
73 &agent_id,
74 MiddlewareData::AfterTool {
75 tool_name: ctx.tool_name.clone(),
76 params: serde_json::Value::Null,
77 result: ctx.result.clone(),
78 },
79 );
80
81 let rt = tokio::runtime::Handle::current();
82 let result = rt.block_on(pipeline.execute(&mw_ctx));
83
84 if matches!(
85 result.action,
86 crate::middleware::MiddlewareAction::Terminate
87 ) {
88 terminate_flag.store(true, Ordering::SeqCst);
89 }
90
91 oxicode_agent::AfterToolCallResult::default()
92 }
93 })
94 as std::boxed::Box<
95 dyn Fn(&oxicode_agent::AfterToolCallContext) -> oxicode_agent::AfterToolCallResult
96 + Send
97 + Sync,
98 >);
99
100 let should_stop_after_turn = Some(Arc::new({
101 let flag = terminate_flag;
102
103 move |_ctx: &oxicode_agent::ShouldStopAfterTurnContext| -> bool {
104 flag.load(Ordering::SeqCst)
105 }
106 })
107 as Arc<dyn Fn(&oxicode_agent::ShouldStopAfterTurnContext) -> bool + Send + Sync>);
108
109 AgentHooks {
110 before_tool_call,
111 after_tool_call,
112 should_stop_after_turn,
113 get_steering_messages: None,
114 get_follow_up_messages: None,
115 tool_execution: oxicode_agent::ToolExecutionMode::Parallel,
116 }
117}
118
119#[cfg(test)]
120mod tests {
121 use super::*;
122
123 #[test]
124 fn test_bridge_returns_valid_hooks() {
125 let pipeline = Arc::new(MiddlewarePipeline::new());
126 let terminate_flag = Arc::new(AtomicBool::new(false));
127
128 let hooks = build_hooks(pipeline, "test-agent".into(), terminate_flag);
129
130 assert!(hooks.before_tool_call.is_some());
131 assert!(hooks.after_tool_call.is_some());
132 assert!(hooks.should_stop_after_turn.is_some());
133 }
134}