Skip to main content

zeph_bench/loaders/tau2_bench/envs/
retail.rs

1// SPDX-FileCopyrightText: 2026 Andrei G <bug-ops>
2// SPDX-License-Identifier: MIT OR Apache-2.0
3
4//! Retail domain environment for tau2-bench.
5//!
6//! Holds the in-memory state loaded from `db.json` and implements [`ToolExecutor`]
7//! so the agent can call all 16 retail tools. Every call is recorded to the shared
8//! [`ActionTrace`] before dispatch.
9
10use std::io::BufReader;
11use std::path::Path;
12use std::sync::{Arc, Mutex};
13
14use serde::Deserialize;
15use tracing::instrument;
16use zeph_common::ToolName;
17use zeph_tools::ToolExecutor;
18use zeph_tools::executor::{ToolCall, ToolError, ToolOutput};
19use zeph_tools::registry::ToolDef;
20
21use crate::error::BenchError;
22use crate::loaders::tau2_bench::data::Action;
23
24use super::{ActionTrace, RecordedToolCall, SnapshotableEnv};
25
26// ─── State types ────────────────────────────────────────────────────────────
27
28#[derive(Debug, Clone, Deserialize, serde::Serialize)]
29struct Address {
30    address1: String,
31    address2: String,
32    city: String,
33    state: String,
34    zip: String,
35    country: String,
36}
37
38#[derive(Debug, Clone, Deserialize, serde::Serialize)]
39struct UserName {
40    first_name: String,
41    last_name: String,
42}
43
44#[derive(Debug, Clone, Deserialize, serde::Serialize)]
45struct RetailUser {
46    user_id: String,
47    name: UserName,
48    email: String,
49    address: Address,
50    payment_methods: serde_json::Map<String, serde_json::Value>,
51    #[serde(default, flatten)]
52    _rest: serde_json::Map<String, serde_json::Value>,
53}
54
55#[derive(Debug, Clone, Deserialize, serde::Serialize)]
56struct OrderItem {
57    item_id: String,
58    name: String,
59    product_id: String,
60    price: f64,
61    #[serde(default)]
62    options: serde_json::Map<String, serde_json::Value>,
63}
64
65#[derive(Debug, Clone, Deserialize, serde::Serialize)]
66struct RetailOrder {
67    order_id: String,
68    user_id: String,
69    address: Address,
70    items: Vec<OrderItem>,
71    status: String,
72    #[serde(default)]
73    payment_history: Vec<serde_json::Value>,
74    #[serde(default, flatten)]
75    _rest: serde_json::Map<String, serde_json::Value>,
76}
77
78/// Full in-memory retail database.
79///
80/// Loaded once from `db.json` via [`RetailState::load`] and then cloned per scenario.
81#[derive(Debug, Clone, Deserialize, serde::Serialize)]
82struct RetailState {
83    /// Product catalogue: `product_id → { name, variants: { item_id → { options, price, available } } }`.
84    products: serde_json::Map<String, serde_json::Value>,
85    /// User records by `user_id`.
86    users: std::collections::HashMap<String, RetailUser>,
87    /// Order records by `order_id`.
88    orders: std::collections::HashMap<String, RetailOrder>,
89}
90
91impl RetailState {
92    fn load(db_path: &Path) -> Result<Self, BenchError> {
93        let file = std::fs::File::open(db_path)
94            .map_err(|e| BenchError::InvalidFormat(format!("open retail db.json: {e}")))?;
95        serde_json::from_reader(BufReader::new(file))
96            .map_err(|e| BenchError::InvalidFormat(format!("parse retail db.json: {e}")))
97    }
98}
99
100// ─── Executor ────────────────────────────────────────────────────────────────
101
102/// In-memory retail environment executor for tau2-bench.
103///
104/// Holds mutable state (orders, users) behind a `std::sync::Mutex` and records every
105/// tool call to the shared [`ActionTrace`] before dispatching.
106///
107/// # Construction
108///
109/// Always use [`RetailEnv::new_from_seed`] — it returns `(Self, ActionTrace)` where
110/// the trace is the same `Arc` the env stores internally.
111pub struct RetailEnv {
112    state: Arc<Mutex<RetailState>>,
113    trace: ActionTrace,
114}
115
116/// Load `RetailState` from `db_path`, memoising the result for the process lifetime.
117///
118/// The cache is keyed by the canonicalized (real) path so different relative paths to the
119/// same file share an entry. Cache entries are never evicted — the process is short-lived
120/// for benchmark runs, so unbounded growth is not a concern.
121///
122/// Lock poisoning falls through to a fresh disk reload; returning a valid state is always
123/// preferable to propagating the poison error.
124fn cached_retail_load(db_path: &Path) -> Result<RetailState, BenchError> {
125    use std::collections::HashMap;
126    use std::sync::{Mutex, OnceLock};
127
128    static CACHE: OnceLock<Mutex<HashMap<std::path::PathBuf, Arc<RetailState>>>> = OnceLock::new();
129    let cache = CACHE.get_or_init(|| Mutex::new(HashMap::new()));
130
131    // Canonicalize so different relative-path spellings share the same entry.
132    let key = std::fs::canonicalize(db_path).unwrap_or_else(|_| db_path.to_path_buf());
133
134    // Fast path: cache hit — clone the Arc (cheap pointer bump) and return.
135    if let Ok(guard) = cache.lock()
136        && let Some(hit) = guard.get(&key)
137    {
138        return Ok((**hit).clone());
139    }
140
141    // Slow path: load from disk, then memoize.
142    let state = RetailState::load(db_path)?;
143    let arc = Arc::new(state.clone());
144    if let Ok(mut guard) = cache.lock() {
145        guard.insert(key, arc);
146    }
147    Ok(state)
148}
149
150impl RetailEnv {
151    /// Load state from `db.json` and return `(env, trace)`.
152    ///
153    /// The returned `ActionTrace` is the same `Arc` stored inside the env. The evaluator
154    /// must hold this clone to read recorded calls after the run completes.
155    ///
156    /// The `db.json` file is loaded once per process per unique path and memoised via
157    /// a process-global cache. Each call receives an independent deep clone of the state
158    /// so mutations in one scenario cannot affect another.
159    ///
160    /// # Errors
161    ///
162    /// Returns [`BenchError::InvalidFormat`] when `db.json` is missing or malformed.
163    pub fn new_from_seed(db_path: &Path) -> Result<(Self, ActionTrace), BenchError> {
164        let state = cached_retail_load(db_path)?;
165        let trace: ActionTrace = Arc::new(Mutex::new(Vec::new()));
166        let env = Self {
167            state: Arc::new(Mutex::new(state)),
168            trace: trace.clone(),
169        };
170        Ok((env, trace))
171    }
172}
173
174impl SnapshotableEnv for RetailEnv {
175    fn state_snapshot(&self) -> serde_json::Value {
176        let state = self.state.lock().expect("state mutex poisoned").clone();
177        serde_json::to_value(&state).unwrap_or(serde_json::Value::Null)
178    }
179}
180
181impl RetailEnv {
182    /// Replay `actions` on this env instance to build an expected database state.
183    ///
184    /// The caller must construct a dedicated fresh [`RetailEnv`] for replay — never call
185    /// this on a post-run env, as it would contaminate the final state with gold actions.
186    ///
187    /// Actions with `requestor != "assistant"` are skipped (they represent user turns).
188    ///
189    /// # Errors
190    ///
191    /// Returns [`BenchError`] if a gold action fails to execute in the env.
192    #[instrument(skip_all, name = "bench.tau2.retail.replay_actions")]
193    pub async fn replay_actions(&self, actions: &[Action]) -> Result<(), BenchError> {
194        for action in actions {
195            if action.requestor != "assistant" {
196                continue;
197            }
198            let call = ToolCall {
199                tool_id: ToolName::new(action.name.as_str()),
200                params: action.arguments.clone(),
201                caller_id: None,
202                context: None,
203                tool_call_id: String::new(),
204                skill_name: None,
205            };
206            self.execute_tool_call(&call).await.map_err(|e| {
207                BenchError::InvalidFormat(format!("replay action '{}': {e}", action.name))
208            })?;
209        }
210        Ok(())
211    }
212}
213
214impl ToolExecutor for RetailEnv {
215    async fn execute(&self, _response: &str) -> Result<Option<ToolOutput>, ToolError> {
216        // tau2-bench uses structured tool calls only, not fenced code blocks.
217        Ok(None)
218    }
219
220    fn tool_definitions(&self) -> Vec<ToolDef> {
221        super::tools::retail_definitions()
222    }
223
224    #[instrument(skip_all, name = "bench.tau2.retail.execute_tool_call")]
225    async fn execute_tool_call(&self, call: &ToolCall) -> Result<Option<ToolOutput>, ToolError> {
226        // Record before dispatching (lock held only for push, never across .await).
227        {
228            let mut t = self.trace.lock().expect("trace mutex poisoned");
229            t.push(RecordedToolCall::from_tool_call(call));
230        }
231
232        let params = &call.params;
233        let summary = match call.tool_id.as_str() {
234            "calculate" => handle_calculate(params)?,
235            "cancel_pending_order" => self.handle_cancel_pending_order(params)?,
236            "exchange_delivered_order_items" => {
237                self.handle_exchange_delivered_order_items(params)?
238            }
239            "find_user_id_by_email" => self.handle_find_user_id_by_email(params)?,
240            "find_user_id_by_name_zip" => self.handle_find_user_id_by_name_zip(params)?,
241            "get_order_details" => self.handle_get_order_details(params)?,
242            "get_product_details" => self.handle_get_product_details(params)?,
243            "get_item_details" => self.handle_get_item_details(params)?,
244            "get_user_details" => self.handle_get_user_details(params)?,
245            "list_all_product_types" => self.handle_list_all_product_types(),
246            "modify_pending_order_address" => self.handle_modify_pending_order_address(params)?,
247            "modify_pending_order_items" => self.handle_modify_pending_order_items(params)?,
248            "modify_pending_order_payment" => self.handle_modify_pending_order_payment(params)?,
249            "modify_user_address" => self.handle_modify_user_address(params)?,
250            "return_delivered_order_items" => self.handle_return_delivered_order_items(params)?,
251            "transfer_to_human_agents" => handle_transfer_to_human_agents(params)?,
252            _ => return Ok(None),
253        };
254
255        Ok(Some(ToolOutput {
256            tool_name: call.tool_id.clone(),
257            summary,
258            blocks_executed: 1,
259            filter_stats: None,
260            diff: None,
261            streamed: false,
262            terminal_id: None,
263            locations: None,
264            raw_response: None,
265            claim_source: None,
266            ..Default::default()
267        }))
268    }
269
270    zeph_tools::tool_executor_no_inner_defaults!();
271}
272
273// ─── Handlers ────────────────────────────────────────────────────────────────
274
275fn params_str<'a>(
276    params: &'a serde_json::Map<String, serde_json::Value>,
277    key: &str,
278) -> Result<&'a str, ToolError> {
279    params
280        .get(key)
281        .and_then(|v| v.as_str())
282        .ok_or_else(|| ToolError::InvalidParams {
283            message: format!("missing or non-string parameter '{key}'"),
284        })
285}
286
287fn params_str_list(
288    params: &serde_json::Map<String, serde_json::Value>,
289    key: &str,
290) -> Result<Vec<String>, ToolError> {
291    params
292        .get(key)
293        .and_then(|v| v.as_array())
294        .ok_or_else(|| ToolError::InvalidParams {
295            message: format!("missing or non-array parameter '{key}'"),
296        })
297        .and_then(|arr| {
298            arr.iter()
299                .map(|v| {
300                    v.as_str()
301                        .map(ToOwned::to_owned)
302                        .ok_or_else(|| ToolError::InvalidParams {
303                            message: format!("array '{key}' contains non-string element"),
304                        })
305                })
306                .collect()
307        })
308}
309
310fn handle_calculate(
311    params: &serde_json::Map<String, serde_json::Value>,
312) -> Result<String, ToolError> {
313    let expr = params_str(params, "expression")?;
314    Ok(format!("expression={expr} result={}", eval_expr(expr)))
315}
316
317fn handle_transfer_to_human_agents(
318    params: &serde_json::Map<String, serde_json::Value>,
319) -> Result<String, ToolError> {
320    let summary = params_str(params, "summary")?;
321    Ok(format!("transferred_to_human=true summary={summary:?}"))
322}
323
324impl RetailEnv {
325    fn handle_cancel_pending_order(
326        &self,
327        params: &serde_json::Map<String, serde_json::Value>,
328    ) -> Result<String, ToolError> {
329        let order_id = params_str(params, "order_id")?;
330        let reason = params_str(params, "reason")?;
331        let mut state = self.state.lock().expect("state mutex poisoned");
332        let order = state
333            .orders
334            .get_mut(order_id)
335            .ok_or_else(|| ToolError::InvalidParams {
336                message: format!("order {order_id} not found"),
337            })?;
338        if order.status != "pending" {
339            return Err(ToolError::InvalidParams {
340                message: format!("order {order_id} is not pending (status={})", order.status),
341            });
342        }
343        "cancelled".clone_into(&mut order.status);
344        Ok(format!(
345            "order_id={order_id} status=cancelled reason={reason}"
346        ))
347    }
348
349    fn handle_exchange_delivered_order_items(
350        &self,
351        params: &serde_json::Map<String, serde_json::Value>,
352    ) -> Result<String, ToolError> {
353        let order_id = params_str(params, "order_id")?;
354        let item_ids = params_str_list(params, "item_ids")?;
355        let new_item_ids = params_str_list(params, "new_item_ids")?;
356        let payment_method_id = params_str(params, "payment_method_id")?;
357        let mut state = self.state.lock().expect("state mutex poisoned");
358        let order = state
359            .orders
360            .get_mut(order_id)
361            .ok_or_else(|| ToolError::InvalidParams {
362                message: format!("order {order_id} not found"),
363            })?;
364        if order.status != "delivered" {
365            return Err(ToolError::InvalidParams {
366                message: format!(
367                    "order {order_id} is not delivered (status={})",
368                    order.status
369                ),
370            });
371        }
372        order.items.retain(|item| !item_ids.contains(&item.item_id));
373        for new_id in &new_item_ids {
374            order.items.push(OrderItem {
375                item_id: new_id.clone(),
376                name: "exchanged_item".into(),
377                product_id: String::new(),
378                price: 0.0,
379                options: serde_json::Map::new(),
380            });
381        }
382        Ok(format!(
383            "order_id={order_id} exchanged={item_ids:?} new_items={new_item_ids:?} payment_method_id={payment_method_id}"
384        ))
385    }
386
387    fn handle_find_user_id_by_email(
388        &self,
389        params: &serde_json::Map<String, serde_json::Value>,
390    ) -> Result<String, ToolError> {
391        let email = params_str(params, "email")?;
392        let state = self.state.lock().expect("state mutex poisoned");
393        let user = state
394            .users
395            .values()
396            .find(|u| u.email.eq_ignore_ascii_case(email))
397            .ok_or_else(|| ToolError::InvalidParams {
398                message: format!("no user found with email {email}"),
399            })?;
400        Ok(format!("user_id={}", user.user_id))
401    }
402
403    fn handle_find_user_id_by_name_zip(
404        &self,
405        params: &serde_json::Map<String, serde_json::Value>,
406    ) -> Result<String, ToolError> {
407        let first = params_str(params, "first_name")?;
408        let last = params_str(params, "last_name")?;
409        let zip = params_str(params, "zip")?;
410        let state = self.state.lock().expect("state mutex poisoned");
411        let user = state
412            .users
413            .values()
414            .find(|u| {
415                u.name.first_name.eq_ignore_ascii_case(first)
416                    && u.name.last_name.eq_ignore_ascii_case(last)
417                    && u.address.zip == zip
418            })
419            .ok_or_else(|| ToolError::InvalidParams {
420                message: format!("no user found for {first} {last} zip={zip}"),
421            })?;
422        Ok(format!("user_id={}", user.user_id))
423    }
424
425    fn handle_get_order_details(
426        &self,
427        params: &serde_json::Map<String, serde_json::Value>,
428    ) -> Result<String, ToolError> {
429        let order_id = params_str(params, "order_id")?;
430        let state = self.state.lock().expect("state mutex poisoned");
431        let order = state
432            .orders
433            .get(order_id)
434            .ok_or_else(|| ToolError::InvalidParams {
435                message: format!("order {order_id} not found"),
436            })?;
437        Ok(serde_json::to_string(order).unwrap_or_else(|_| format!("order_id={order_id}")))
438    }
439
440    fn handle_get_product_details(
441        &self,
442        params: &serde_json::Map<String, serde_json::Value>,
443    ) -> Result<String, ToolError> {
444        let product_id = params_str(params, "product_id")?;
445        let state = self.state.lock().expect("state mutex poisoned");
446        let product = state
447            .products
448            .get(product_id)
449            .ok_or_else(|| ToolError::InvalidParams {
450                message: format!("product {product_id} not found"),
451            })?;
452        Ok(product.to_string())
453    }
454
455    fn handle_get_item_details(
456        &self,
457        params: &serde_json::Map<String, serde_json::Value>,
458    ) -> Result<String, ToolError> {
459        let item_id = params_str(params, "item_id")?;
460        let state = self.state.lock().expect("state mutex poisoned");
461        // Items (variants) are nested inside each product's `variants` map.
462        for product in state.products.values() {
463            if let Some(variants) = product.get("variants").and_then(|v| v.as_object())
464                && let Some(variant) = variants.get(item_id)
465            {
466                return Ok(variant.to_string());
467            }
468        }
469        Err(ToolError::InvalidParams {
470            message: format!("item {item_id} not found"),
471        })
472    }
473
474    fn handle_get_user_details(
475        &self,
476        params: &serde_json::Map<String, serde_json::Value>,
477    ) -> Result<String, ToolError> {
478        let user_id = params_str(params, "user_id")?;
479        let state = self.state.lock().expect("state mutex poisoned");
480        let user = state
481            .users
482            .get(user_id)
483            .ok_or_else(|| ToolError::InvalidParams {
484                message: format!("user {user_id} not found"),
485            })?;
486        Ok(serde_json::to_string(user).unwrap_or_else(|_| format!("user_id={user_id}")))
487    }
488
489    fn handle_list_all_product_types(&self) -> String {
490        let state = self.state.lock().expect("state mutex poisoned");
491        let names: Vec<&str> = state
492            .products
493            .values()
494            .filter_map(|p| p.get("name").and_then(|n| n.as_str()))
495            .collect();
496        format!("product_types={names:?}")
497    }
498
499    fn handle_modify_pending_order_address(
500        &self,
501        params: &serde_json::Map<String, serde_json::Value>,
502    ) -> Result<String, ToolError> {
503        let order_id = params_str(params, "order_id")?;
504        let address1 = params_str(params, "address1")?.to_owned();
505        let address2 = params_str(params, "address2")?.to_owned();
506        let city = params_str(params, "city")?.to_owned();
507        let state_str = params_str(params, "state")?.to_owned();
508        let zip = params_str(params, "zip")?.to_owned();
509        let country = params_str(params, "country")?.to_owned();
510        let mut state = self.state.lock().expect("state mutex poisoned");
511        let order = state
512            .orders
513            .get_mut(order_id)
514            .ok_or_else(|| ToolError::InvalidParams {
515                message: format!("order {order_id} not found"),
516            })?;
517        if order.status != "pending" {
518            return Err(ToolError::InvalidParams {
519                message: format!("order {order_id} is not pending (status={})", order.status),
520            });
521        }
522        order.address = Address {
523            address1,
524            address2,
525            city,
526            state: state_str,
527            zip,
528            country,
529        };
530        Ok(format!("order_id={order_id} address_updated=true"))
531    }
532
533    fn handle_modify_pending_order_items(
534        &self,
535        params: &serde_json::Map<String, serde_json::Value>,
536    ) -> Result<String, ToolError> {
537        let order_id = params_str(params, "order_id")?;
538        let item_ids = params_str_list(params, "item_ids")?;
539        let new_item_ids = params_str_list(params, "new_item_ids")?;
540        let payment_method_id = params_str(params, "payment_method_id")?;
541        let mut state = self.state.lock().expect("state mutex poisoned");
542        let order = state
543            .orders
544            .get_mut(order_id)
545            .ok_or_else(|| ToolError::InvalidParams {
546                message: format!("order {order_id} not found"),
547            })?;
548        if order.status != "pending" {
549            return Err(ToolError::InvalidParams {
550                message: format!("order {order_id} is not pending (status={})", order.status),
551            });
552        }
553        order.items.retain(|item| !item_ids.contains(&item.item_id));
554        for new_id in &new_item_ids {
555            order.items.push(OrderItem {
556                item_id: new_id.clone(),
557                name: "new_item".into(),
558                product_id: String::new(),
559                price: 0.0,
560                options: serde_json::Map::new(),
561            });
562        }
563        Ok(format!(
564            "order_id={order_id} removed={item_ids:?} added={new_item_ids:?} payment_method_id={payment_method_id}"
565        ))
566    }
567
568    fn handle_modify_pending_order_payment(
569        &self,
570        params: &serde_json::Map<String, serde_json::Value>,
571    ) -> Result<String, ToolError> {
572        let order_id = params_str(params, "order_id")?;
573        let payment_method_id = params_str(params, "payment_method_id")?;
574        let state = self.state.lock().expect("state mutex poisoned");
575        let order = state
576            .orders
577            .get(order_id)
578            .ok_or_else(|| ToolError::InvalidParams {
579                message: format!("order {order_id} not found"),
580            })?;
581        if order.status != "pending" {
582            return Err(ToolError::InvalidParams {
583                message: format!("order {order_id} is not pending (status={})", order.status),
584            });
585        }
586        Ok(format!(
587            "order_id={order_id} payment_method_id={payment_method_id} updated=true"
588        ))
589    }
590
591    fn handle_modify_user_address(
592        &self,
593        params: &serde_json::Map<String, serde_json::Value>,
594    ) -> Result<String, ToolError> {
595        let user_id = params_str(params, "user_id")?;
596        let address1 = params_str(params, "address1")?.to_owned();
597        let address2 = params_str(params, "address2")?.to_owned();
598        let city = params_str(params, "city")?.to_owned();
599        let state_str = params_str(params, "state")?.to_owned();
600        let zip = params_str(params, "zip")?.to_owned();
601        let country = params_str(params, "country")?.to_owned();
602        let mut state = self.state.lock().expect("state mutex poisoned");
603        let user = state
604            .users
605            .get_mut(user_id)
606            .ok_or_else(|| ToolError::InvalidParams {
607                message: format!("user {user_id} not found"),
608            })?;
609        user.address = Address {
610            address1,
611            address2,
612            city,
613            state: state_str,
614            zip,
615            country,
616        };
617        Ok(format!("user_id={user_id} address_updated=true"))
618    }
619
620    fn handle_return_delivered_order_items(
621        &self,
622        params: &serde_json::Map<String, serde_json::Value>,
623    ) -> Result<String, ToolError> {
624        let order_id = params_str(params, "order_id")?;
625        let item_ids = params_str_list(params, "item_ids")?;
626        let payment_method_id = params_str(params, "payment_method_id")?;
627        let mut state = self.state.lock().expect("state mutex poisoned");
628        let order = state
629            .orders
630            .get_mut(order_id)
631            .ok_or_else(|| ToolError::InvalidParams {
632                message: format!("order {order_id} not found"),
633            })?;
634        if order.status != "delivered" {
635            return Err(ToolError::InvalidParams {
636                message: format!(
637                    "order {order_id} is not delivered (status={})",
638                    order.status
639                ),
640            });
641        }
642        let refund: f64 = order
643            .items
644            .iter()
645            .filter(|i| item_ids.contains(&i.item_id))
646            .map(|i| i.price)
647            .sum();
648        order.items.retain(|i| !item_ids.contains(&i.item_id));
649        Ok(format!(
650            "order_id={order_id} returned={item_ids:?} refund={refund:.2} payment_method_id={payment_method_id}"
651        ))
652    }
653}
654
655/// Minimal expression evaluator for `calculate` (handles +, -, *, / on f64).
656fn eval_expr(expr: &str) -> String {
657    // Simple left-to-right evaluation without precedence for MVP.
658    // Handles: "1 + 2", "10 * 3.5", "100 / 4 - 5".
659    let tokens: Vec<&str> = expr.split_whitespace().collect();
660    if tokens.is_empty() {
661        return "NaN".into();
662    }
663    let mut result: f64 = match tokens[0].parse() {
664        Ok(v) => v,
665        Err(_) => return "NaN".into(),
666    };
667    let mut i = 1;
668    while i + 1 < tokens.len() {
669        let op = tokens[i];
670        let right: f64 = match tokens[i + 1].parse() {
671            Ok(v) => v,
672            Err(_) => return "NaN".into(),
673        };
674        match op {
675            "+" => result += right,
676            "-" => result -= right,
677            "*" => result *= right,
678            "/" => {
679                if right == 0.0 {
680                    return "division by zero".into();
681                }
682                result /= right;
683            }
684            _ => return format!("unknown op: {op}"),
685        }
686        i += 2;
687    }
688    format!("{result}")
689}
690
691#[cfg(test)]
692mod tests {
693    use super::*;
694    use std::assert_matches;
695
696    const RETAIL_DB_MIN: &str = r##"{
697        "products": {
698            "prod_001": {
699                "name": "T-Shirt",
700                "product_id": "prod_001",
701                "variants": {
702                    "item_001": {
703                        "item_id": "item_001",
704                        "options": {"color": "blue", "size": "M"},
705                        "available": true,
706                        "price": 25.00
707                    }
708                }
709            }
710        },
711        "users": {
712            "alice_smith_1": {
713                "user_id": "alice_smith_1",
714                "name": {"first_name": "Alice", "last_name": "Smith"},
715                "email": "alice@example.com",
716                "address": {
717                    "address1": "1 Main St",
718                    "address2": "",
719                    "city": "Boston",
720                    "state": "MA",
721                    "zip": "02101",
722                    "country": "USA"
723                },
724                "payment_methods": {
725                    "credit_card_1": {"source": "credit_card", "id": "credit_card_1"}
726                }
727            }
728        },
729        "orders": {
730            "#W0001": {
731                "order_id": "#W0001",
732                "user_id": "alice_smith_1",
733                "address": {
734                    "address1": "1 Main St",
735                    "address2": "",
736                    "city": "Boston",
737                    "state": "MA",
738                    "zip": "02101",
739                    "country": "USA"
740                },
741                "items": [
742                    {
743                        "item_id": "item_001",
744                        "name": "T-Shirt",
745                        "product_id": "prod_001",
746                        "price": 25.00,
747                        "options": {"color": "blue", "size": "M"}
748                    }
749                ],
750                "status": "pending",
751                "payment_history": [
752                    {"transaction_type": "payment", "amount": 25.00, "payment_method_id": "credit_card_1"}
753                ]
754            },
755            "#W0002": {
756                "order_id": "#W0002",
757                "user_id": "alice_smith_1",
758                "address": {
759                    "address1": "1 Main St",
760                    "address2": "",
761                    "city": "Boston",
762                    "state": "MA",
763                    "zip": "02101",
764                    "country": "USA"
765                },
766                "items": [
767                    {
768                        "item_id": "item_001",
769                        "name": "T-Shirt",
770                        "product_id": "prod_001",
771                        "price": 25.00,
772                        "options": {"color": "blue", "size": "M"}
773                    }
774                ],
775                "status": "delivered",
776                "payment_history": []
777            }
778        }
779    }"##;
780
781    fn make_env() -> (RetailEnv, ActionTrace) {
782        let dir = tempfile::tempdir().unwrap();
783        let db_path = dir.path().join("db.json");
784        std::fs::write(&db_path, RETAIL_DB_MIN).unwrap();
785        // Keep the tempdir alive by leaking it — acceptable in tests.
786        std::mem::forget(dir);
787        RetailEnv::new_from_seed(&db_path).unwrap()
788    }
789
790    #[allow(clippy::needless_pass_by_value)]
791    fn call(tool: &str, params: serde_json::Value) -> ToolCall {
792        use zeph_common::ToolName;
793        ToolCall {
794            tool_id: ToolName::new(tool),
795            params: params.as_object().cloned().unwrap_or_default(),
796            caller_id: None,
797            context: None,
798
799            tool_call_id: String::new(),
800            skill_name: None,
801        }
802    }
803
804    #[tokio::test]
805    async fn find_user_by_email() {
806        let (env, _) = make_env();
807        let c = call(
808            "find_user_id_by_email",
809            serde_json::json!({"email": "alice@example.com"}),
810        );
811        let out = env.execute_tool_call(&c).await.unwrap().unwrap();
812        assert!(out.summary.contains("alice_smith_1"));
813    }
814
815    #[tokio::test]
816    async fn find_user_by_name_zip() {
817        let (env, _) = make_env();
818        let c = call(
819            "find_user_id_by_name_zip",
820            serde_json::json!({"first_name": "Alice", "last_name": "Smith", "zip": "02101"}),
821        );
822        let out = env.execute_tool_call(&c).await.unwrap().unwrap();
823        assert!(out.summary.contains("alice_smith_1"));
824    }
825
826    #[tokio::test]
827    async fn cancel_pending_order_success() {
828        let (env, trace) = make_env();
829        let c = call(
830            "cancel_pending_order",
831            serde_json::json!({"order_id": "#W0001", "reason": "no_longer_needed"}),
832        );
833        let out = env.execute_tool_call(&c).await.unwrap().unwrap();
834        assert!(out.summary.contains("cancelled"));
835        assert_eq!(trace.lock().unwrap().len(), 1);
836    }
837
838    #[tokio::test]
839    async fn cancel_non_pending_order_fails() {
840        let (env, _) = make_env();
841        let c = call(
842            "cancel_pending_order",
843            serde_json::json!({"order_id": "#W0002", "reason": "no_longer_needed"}),
844        );
845        let err = env.execute_tool_call(&c).await.unwrap_err();
846        assert_matches!(err, ToolError::InvalidParams { .. });
847    }
848
849    #[tokio::test]
850    async fn get_order_details_success() {
851        let (env, _) = make_env();
852        let c = call(
853            "get_order_details",
854            serde_json::json!({"order_id": "#W0001"}),
855        );
856        let out = env.execute_tool_call(&c).await.unwrap().unwrap();
857        assert!(out.summary.contains("W0001") || out.summary.contains("pending"));
858    }
859
860    #[tokio::test]
861    async fn trace_records_calls() {
862        let (env, trace) = make_env();
863        assert_eq!(Arc::strong_count(&trace), 2, "env must share the trace Arc");
864        let c = call(
865            "get_user_details",
866            serde_json::json!({"user_id": "alice_smith_1"}),
867        );
868        let _ = env.execute_tool_call(&c).await;
869        assert_eq!(trace.lock().unwrap().len(), 1);
870        assert_eq!(trace.lock().unwrap()[0].name, "get_user_details");
871    }
872
873    #[test]
874    fn eval_expr_add() {
875        assert_eq!(eval_expr("1 + 2"), "3");
876    }
877
878    #[test]
879    fn eval_expr_multiply() {
880        assert_eq!(eval_expr("3 * 4"), "12");
881    }
882
883    #[test]
884    fn eval_expr_divide_by_zero() {
885        assert!(eval_expr("1 / 0").contains("zero"));
886    }
887
888    /// `new_from_seed` called twice with the same path must return independently mutable
889    /// environments (mutations in one must not affect the other).
890    #[tokio::test]
891    async fn new_from_seed_returns_independent_copies() {
892        let dir = tempfile::tempdir().unwrap();
893        let db_path = dir.path().join("db_iso.json");
894        std::fs::write(&db_path, RETAIL_DB_MIN).unwrap();
895
896        let (env1, _trace1) = RetailEnv::new_from_seed(&db_path).unwrap();
897        let (env2, _trace2) = RetailEnv::new_from_seed(&db_path).unwrap();
898
899        // Cancel the pending order #W0001 in env1.
900        let c = call(
901            "cancel_pending_order",
902            serde_json::json!({"order_id": "#W0001", "reason": "changed mind"}),
903        );
904        env1.execute_tool_call(&c).await.unwrap();
905
906        // env2 must still have the order.
907        let get = call(
908            "get_order_details",
909            serde_json::json!({"order_id": "#W0001"}),
910        );
911        assert!(
912            env2.execute_tool_call(&get).await.unwrap().is_some(),
913            "mutation in env1 must not affect env2"
914        );
915    }
916}