Skip to main content

roder_core/
fake_provider.rs

1use futures::stream;
2use roder_api::catalog::{PROVIDER_MOCK, models_for_provider};
3use roder_api::extension::InferenceEngineId;
4use roder_api::inference::*;
5use roder_api::transcript::TranscriptItem;
6
7mod tbench_diagnostics;
8
9pub struct FakeInferenceEngine;
10
11#[async_trait::async_trait]
12impl InferenceEngine for FakeInferenceEngine {
13    fn id(&self) -> InferenceEngineId {
14        PROVIDER_MOCK.to_string()
15    }
16
17    fn capabilities(&self) -> InferenceCapabilities {
18        InferenceCapabilities::text_only()
19    }
20
21    async fn list_models(
22        &self,
23        _ctx: InferenceProviderContext<'_>,
24    ) -> anyhow::Result<Vec<ModelDescriptor>> {
25        Ok(models_for_provider(PROVIDER_MOCK, true))
26    }
27
28    async fn stream_turn(
29        &self,
30        _ctx: InferenceTurnContext<'_>,
31        request: AgentInferenceRequest,
32    ) -> anyhow::Result<InferenceEventStream> {
33        if should_request_user_input(&request) {
34            let stream = stream::iter(vec![Ok(InferenceEvent::ToolCallCompleted(
35                ToolCallCompleted {
36                    id: "fake-user-input".to_string(),
37                    name: "request_user_input".to_string(),
38                    arguments: serde_json::json!({
39                        "questions": [{
40                            "header": "Choice",
41                            "id": "choice",
42                            "question": "Which option should be used?",
43                            "options": [
44                                { "label": "A", "description": "Use option A." },
45                                { "label": "B", "description": "Use option B." }
46                            ]
47                        }]
48                    })
49                    .to_string(),
50                },
51            ))]);
52            return Ok(Box::pin(stream));
53        }
54        if should_call_external_tool(&request) {
55            let stream = stream::iter(vec![Ok(InferenceEvent::ToolCallCompleted(
56                ToolCallCompleted {
57                    id: "fake-external-tool".to_string(),
58                    name: "acme_lookup".to_string(),
59                    arguments: serde_json::json!({ "query": "thread status" }).to_string(),
60                },
61            ))]);
62            return Ok(Box::pin(stream));
63        }
64        if should_update_task_ledger(&request) {
65            let complete = prompt_contains(&request, "FAKE_TASK_LEDGER_COMPLETE");
66            let stream = stream::iter(vec![Ok(InferenceEvent::ToolCallCompleted(
67                ToolCallCompleted {
68                    id: "fake-task-ledger".to_string(),
69                    name: "task_ledger.update".to_string(),
70                    arguments: task_ledger_arguments(complete),
71                },
72            ))]);
73            return Ok(Box::pin(stream));
74        }
75        if let Some(tool_call) = tbench_diagnostics::next_tool_call(&request) {
76            let stream = stream::iter(vec![Ok(InferenceEvent::ToolCallCompleted(tool_call))]);
77            return Ok(Box::pin(stream));
78        }
79        if should_write_file(&request) {
80            let stream = stream::iter(vec![Ok(InferenceEvent::ToolCallCompleted(
81                ToolCallCompleted {
82                    id: "fake-write-file".to_string(),
83                    name: "write_file".to_string(),
84                    arguments: serde_json::json!({
85                        "path": "src/lib.rs",
86                        "content": "pub fn fake() -> &'static str { \"verified\" }\n"
87                    })
88                    .to_string(),
89                },
90            ))]);
91            return Ok(Box::pin(stream));
92        }
93        if should_grep(&request) {
94            let stream = stream::iter(vec![Ok(InferenceEvent::ToolCallCompleted(
95                ToolCallCompleted {
96                    id: "fake-grep".to_string(),
97                    name: "grep".to_string(),
98                    arguments: serde_json::json!({
99                        "query": "BUG_ROOT_CAUSE_TOKEN",
100                        "path": ".",
101                        "mode": "indexed",
102                        "limit": 20
103                    })
104                    .to_string(),
105                },
106            ))]);
107            return Ok(Box::pin(stream));
108        }
109        if should_zerolang_graph_dump(&request) {
110            let stream = stream::iter(vec![Ok(InferenceEvent::ToolCallCompleted(
111                ToolCallCompleted {
112                    id: "fake-zerolang-graph-dump".to_string(),
113                    name: "zerolang_graph_dump".to_string(),
114                    arguments: serde_json::json!({
115                        "input": "src/main.0"
116                    })
117                    .to_string(),
118                },
119            ))]);
120            return Ok(Box::pin(stream));
121        }
122        if should_zerolang_edit(&request) {
123            let stream = stream::iter(vec![Ok(InferenceEvent::ToolCallCompleted(
124                ToolCallCompleted {
125                    id: "fake-zerolang-edit".to_string(),
126                    name: "zerolang_edit".to_string(),
127                    arguments: serde_json::json!({
128                        "input": "src/main.0",
129                        "graphHash": "graph:f76987e99677f1b3",
130                        "operations": [{
131                            "op": "set",
132                            "node": "#610c78bf",
133                            "field": "value",
134                            "expect": "hello from zero\n",
135                            "value": "hello from roder\n"
136                        }],
137                        "validate": true
138                    })
139                    .to_string(),
140                },
141            ))]);
142            return Ok(Box::pin(stream));
143        }
144        if should_discovery_read(&request) {
145            let stream = stream::iter(vec![Ok(InferenceEvent::ToolCallCompleted(
146                ToolCallCompleted {
147                    id: "fake-discovery-read".to_string(),
148                    name: "discovery.read".to_string(),
149                    arguments: serde_json::json!({
150                        "item_id": "tool:builtin-coding-tools/grep",
151                        "promote": true,
152                        "limit": 20
153                    })
154                    .to_string(),
155                },
156            ))]);
157            return Ok(Box::pin(stream));
158        }
159        if should_discovery_search(&request) {
160            let stream = stream::iter(vec![Ok(InferenceEvent::ToolCallCompleted(
161                ToolCallCompleted {
162                    id: "fake-discovery-search".to_string(),
163                    name: "discovery.search".to_string(),
164                    arguments: serde_json::json!({
165                        "query": "grep",
166                        "limit": 20
167                    })
168                    .to_string(),
169                },
170            ))]);
171            return Ok(Box::pin(stream));
172        }
173        if should_spawn_fake_agent(&request) {
174            let stream = stream::iter(vec![Ok(InferenceEvent::ToolCallCompleted(
175                ToolCallCompleted {
176                    id: "fake-spawn-agent".to_string(),
177                    name: "spawn_agent".to_string(),
178                    arguments: serde_json::json!({
179                        "task_name": "reviewer",
180                        "message": "review the fake agent control smoke"
181                    })
182                    .to_string(),
183                },
184            ))]);
185            return Ok(Box::pin(stream));
186        }
187        if should_list_fake_agents(&request) {
188            let stream = stream::iter(vec![Ok(InferenceEvent::ToolCallCompleted(
189                ToolCallCompleted {
190                    id: "fake-list-agents".to_string(),
191                    name: "list_agents".to_string(),
192                    arguments: "{}".to_string(),
193                },
194            ))]);
195            return Ok(Box::pin(stream));
196        }
197        if should_message_fake_agent(&request) {
198            let stream = stream::iter(vec![Ok(InferenceEvent::ToolCallCompleted(
199                ToolCallCompleted {
200                    id: "fake-send-message".to_string(),
201                    name: "send_message".to_string(),
202                    arguments: serde_json::json!({
203                        "target": "reviewer",
204                        "message": "add one more fake smoke detail"
205                    })
206                    .to_string(),
207                },
208            ))]);
209            return Ok(Box::pin(stream));
210        }
211        if should_wait_fake_agent(&request) {
212            let stream = stream::iter(vec![Ok(InferenceEvent::ToolCallCompleted(
213                ToolCallCompleted {
214                    id: "fake-wait-agent".to_string(),
215                    name: "wait_agent".to_string(),
216                    arguments: serde_json::json!({
217                        "target": "reviewer",
218                        "timeout_ms": 1000
219                    })
220                    .to_string(),
221                },
222            ))]);
223            return Ok(Box::pin(stream));
224        }
225        if should_close_fake_agent(&request) {
226            let stream = stream::iter(vec![Ok(InferenceEvent::ToolCallCompleted(
227                ToolCallCompleted {
228                    id: "fake-close-agent".to_string(),
229                    name: "close_agent".to_string(),
230                    arguments: serde_json::json!({
231                        "target": "reviewer"
232                    })
233                    .to_string(),
234                },
235            ))]);
236            return Ok(Box::pin(stream));
237        }
238        if should_complete_verification(&request) {
239            let failed = prompt_contains(&request, "FAKE_VERIFICATION_FAILED");
240            let stream = stream::iter(vec![Ok(InferenceEvent::ToolCallCompleted(
241                ToolCallCompleted {
242                    id: "fake-verification".to_string(),
243                    name: "verification_review".to_string(),
244                    arguments: verification_arguments(failed),
245                },
246            ))]);
247            return Ok(Box::pin(stream));
248        }
249        if verification_failed(&request) {
250            let stream = stream::iter(vec![Ok(InferenceEvent::Failed(InferenceFailure {
251                message: "verification gaps remain: tests not run".to_string(),
252            }))]);
253            return Ok(Box::pin(stream));
254        }
255        if user_input_unavailable(&request) {
256            let stream = stream::iter(vec![Ok(InferenceEvent::Failed(InferenceFailure {
257                message: "clarification unavailable in non-interactive runtime profile".to_string(),
258            }))]);
259            return Ok(Box::pin(stream));
260        }
261        let stream = stream::iter(vec![
262            Ok(InferenceEvent::MessageDelta(MessageDelta {
263                text: "hello".to_string(),
264                phase: None,
265            })),
266            Ok(InferenceEvent::MessageDelta(MessageDelta {
267                text: " from".to_string(),
268                phase: None,
269            })),
270            Ok(InferenceEvent::MessageDelta(MessageDelta {
271                text: " roder".to_string(),
272                phase: None,
273            })),
274            Ok(InferenceEvent::Completed(CompletionMetadata {
275                stop_reason: Some("stop".to_string()),
276                provider_response_id: None,
277            })),
278        ]);
279
280        Ok(Box::pin(stream))
281    }
282}
283
284fn should_request_user_input(request: &AgentInferenceRequest) -> bool {
285    prompt_contains(request, "FAKE_REQUEST_USER_INPUT")
286        && !request.transcript.iter().any(|item| {
287            matches!(
288                item,
289                TranscriptItem::ToolResult(result)
290                    if result.name.as_deref() == Some("request_user_input")
291            )
292        })
293}
294
295fn user_input_unavailable(request: &AgentInferenceRequest) -> bool {
296    request.transcript.iter().any(|item| {
297        matches!(
298            item,
299            TranscriptItem::ToolResult(result)
300                if result.name.as_deref() == Some("request_user_input")
301                    && result.is_error
302                    && result.result.contains("User input is unavailable")
303        )
304    })
305}
306
307fn should_call_external_tool(request: &AgentInferenceRequest) -> bool {
308    prompt_contains(request, "FAKE_EXTERNAL_TOOL") && !has_tool_result(request, "acme_lookup")
309}
310
311fn should_update_task_ledger(request: &AgentInferenceRequest) -> bool {
312    (prompt_contains(request, "FAKE_TASK_LEDGER_UPDATE")
313        || prompt_contains(request, "FAKE_TASK_LEDGER_COMPLETE"))
314        && !request.transcript.iter().any(|item| {
315            matches!(
316                item,
317                TranscriptItem::ToolResult(result)
318                    if result.name.as_deref() == Some("task_ledger.update")
319            )
320        })
321}
322
323fn should_write_file(request: &AgentInferenceRequest) -> bool {
324    prompt_contains(request, "FAKE_WRITE_FILE")
325        && !request.transcript.iter().any(|item| {
326            matches!(
327                item,
328                TranscriptItem::ToolResult(result)
329                    if result.name.as_deref() == Some("write_file")
330            )
331        })
332}
333
334fn should_grep(request: &AgentInferenceRequest) -> bool {
335    prompt_contains(request, "FAKE_GREP_INDEXED")
336        && !request.transcript.iter().any(|item| {
337            matches!(
338                item,
339                TranscriptItem::ToolResult(result) if result.name.as_deref() == Some("grep")
340            )
341        })
342}
343
344fn should_zerolang_graph_dump(request: &AgentInferenceRequest) -> bool {
345    prompt_contains(request, "FAKE_ZEROLANG_GRAPH_EDIT")
346        && !has_tool_result(request, "zerolang_graph_dump")
347}
348
349fn should_zerolang_edit(request: &AgentInferenceRequest) -> bool {
350    prompt_contains(request, "FAKE_ZEROLANG_GRAPH_EDIT")
351        && has_tool_result(request, "zerolang_graph_dump")
352        && !has_tool_result(request, "zerolang_edit")
353}
354
355fn should_discovery_search(request: &AgentInferenceRequest) -> bool {
356    prompt_contains(request, "FAKE_DISCOVERY_SEARCH")
357        && !request.transcript.iter().any(|item| {
358            matches!(
359                item,
360                TranscriptItem::ToolResult(result)
361                    if result.name.as_deref() == Some("discovery.search")
362            )
363        })
364}
365
366fn should_discovery_read(request: &AgentInferenceRequest) -> bool {
367    prompt_contains(request, "FAKE_DISCOVERY_PROMOTE")
368        && request.transcript.iter().any(|item| {
369            matches!(
370                item,
371                TranscriptItem::ToolResult(result)
372                    if result.name.as_deref() == Some("discovery.search")
373            )
374        })
375        && !request.transcript.iter().any(|item| {
376            matches!(
377                item,
378                TranscriptItem::ToolResult(result)
379                    if result.name.as_deref() == Some("discovery.read")
380            )
381        })
382}
383
384fn should_spawn_fake_agent(request: &AgentInferenceRequest) -> bool {
385    prompt_contains(request, "FAKE_AGENT_CONTROL_SMOKE") && !has_tool_result(request, "spawn_agent")
386}
387
388fn should_list_fake_agents(request: &AgentInferenceRequest) -> bool {
389    prompt_contains(request, "FAKE_AGENT_CONTROL_SMOKE")
390        && has_tool_result(request, "spawn_agent")
391        && !has_tool_result(request, "list_agents")
392}
393
394fn should_message_fake_agent(request: &AgentInferenceRequest) -> bool {
395    prompt_contains(request, "FAKE_AGENT_CONTROL_SMOKE")
396        && has_tool_result(request, "list_agents")
397        && !has_tool_result(request, "send_message")
398}
399
400fn should_wait_fake_agent(request: &AgentInferenceRequest) -> bool {
401    prompt_contains(request, "FAKE_AGENT_CONTROL_SMOKE")
402        && has_tool_result(request, "send_message")
403        && !has_tool_result(request, "wait_agent")
404}
405
406fn should_close_fake_agent(request: &AgentInferenceRequest) -> bool {
407    prompt_contains(request, "FAKE_AGENT_CONTROL_SMOKE")
408        && has_tool_result(request, "wait_agent")
409        && !has_tool_result(request, "close_agent")
410}
411
412fn should_complete_verification(request: &AgentInferenceRequest) -> bool {
413    request.transcript.iter().any(|item| {
414        matches!(
415            item,
416            TranscriptItem::UserMessage(message)
417                if message.text.contains("Verification gate blocked final completion")
418        )
419    }) && !request.transcript.iter().any(|item| {
420        matches!(
421            item,
422            TranscriptItem::ToolResult(result)
423                if result.name.as_deref() == Some("verification_review")
424        )
425    })
426}
427
428fn verification_failed(request: &AgentInferenceRequest) -> bool {
429    request.transcript.iter().any(|item| {
430        matches!(
431            item,
432            TranscriptItem::ToolResult(result)
433                if result.name.as_deref() == Some("verification_review")
434                    && result.result.contains("Verification failed")
435        )
436    })
437}
438
439fn has_tool_result(request: &AgentInferenceRequest, name: &str) -> bool {
440    request.transcript.iter().any(|item| {
441        matches!(
442            item,
443            TranscriptItem::ToolResult(result) if result.name.as_deref() == Some(name)
444        )
445    })
446}
447
448fn prompt_contains(request: &AgentInferenceRequest, needle: &str) -> bool {
449    request.transcript.iter().any(|item| {
450        matches!(
451            item,
452            TranscriptItem::UserMessage(message) if message.text.contains(needle)
453        )
454    })
455}
456
457fn task_ledger_arguments(complete: bool) -> String {
458    let second_status = if complete { "completed" } else { "in_progress" };
459    let mut second = serde_json::json!({
460        "id": "verify",
461        "content": "Verify task",
462        "status": second_status
463    });
464    if complete {
465        second["evidence"] = serde_json::json!("fake-provider");
466    }
467    serde_json::json!({
468        "tasks": [
469            { "id": "inspect", "content": "Inspect task", "status": "completed", "evidence": "fake-provider" },
470            second
471        ],
472        "requireCompletionEvidence": true
473    })
474    .to_string()
475}
476
477fn verification_arguments(failed: bool) -> String {
478    let (status, open_gaps) = if failed {
479        ("failed", serde_json::json!(["tests not run"]))
480    } else {
481        ("completed", serde_json::json!([]))
482    };
483    serde_json::json!({
484        "originalTask": "fake verification eval",
485        "changedFiles": ["src/lib.rs"],
486        "toolEvidence": ["write_file wrote src/lib.rs"],
487        "testsRun": if failed { serde_json::json!([]) } else { serde_json::json!(["cargo test -p roder-evals verification"]) },
488        "openGaps": open_gaps,
489        "status": status
490    })
491    .to_string()
492}