Skip to main content

oxicode_sdk/middleware/
bridge.rs

1// MiddlewareBridge — converts MiddlewarePipeline to AgentHooks
2
3use std::sync::Arc;
4use std::sync::atomic::{AtomicBool, Ordering};
5
6use oxicode_agent::AgentHooks;
7
8use crate::middleware::{MiddlewareContext, MiddlewareData, MiddlewarePhase, MiddlewarePipeline};
9
10/// Build `AgentHooks` from a `MiddlewarePipeline`.
11///
12/// Maps MiddlewarePhase to the agent's BeforeTool and AfterTool hooks.
13/// When a middleware calls MiddlewareResult::terminate, sets `terminate_flag`.
14pub 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            // Use block_on to run the async pipeline synchronously from sync callback
35            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}