1#![allow(dead_code)]
2use crate::error::ScriptError;
15use crate::sdk::SdkContext;
16use luft_core::contract::backend::{AgentStatus, RunContext};
17use luft_core::contract::event::AgentEvent;
18use luft_core::contract::finding::Finding;
19use luft_core::Scheduler;
20use mlua::{Lua, Table, Value};
21use serde::{Deserialize, Serialize};
22use std::sync::atomic::Ordering;
23use std::sync::Arc;
24
25#[derive(Debug, Clone, Serialize, Deserialize)]
27pub struct ConvergeConfig {
28 pub adversarial: bool,
30 pub vote_threshold: f32,
32 pub max_rounds: u32,
34 pub producers_per_item: u32,
36 pub adversaries_per_finding: u32,
38 pub model: Option<String>,
40}
41
42impl Default for ConvergeConfig {
43 fn default() -> Self {
44 Self {
45 adversarial: true,
46 vote_threshold: 0.7,
47 max_rounds: 3,
48 producers_per_item: 1,
49 adversaries_per_finding: 1,
50 model: None,
51 }
52 }
53}
54
55#[derive(Debug, Clone, Serialize, Deserialize)]
57pub struct ConvergeResult {
58 pub surviving_items: Vec<serde_json::Value>,
60 pub findings: Vec<Finding>,
62 pub rounds: u32,
64 pub converged: bool,
66 pub round_stats: Vec<RoundStats>,
68}
69
70#[derive(Debug, Clone, Serialize, Deserialize)]
72pub struct RoundStats {
73 pub round: u32,
74 pub items_input: usize,
75 pub findings_generated: usize,
76 pub findings_survived: usize,
77 pub findings_refuted: usize,
78 pub approval_rate: f32,
79}
80
81struct ConvergeState {
83 items: Vec<serde_json::Value>,
84 findings: Vec<Finding>,
85 round_stats: Vec<RoundStats>,
86 converged: bool,
87}
88
89pub async fn execute_convergence(
97 items: Vec<serde_json::Value>,
98 producer_prompt: &str,
99 adversary_prompt: &str,
100 config: ConvergeConfig,
101 scheduler: &Arc<Scheduler>,
102 run_ctx: &RunContext,
103) -> Result<ConvergeResult, ScriptError> {
104 if items.is_empty() {
105 tracing::debug!("converge: no items, returning empty");
106 return Ok(ConvergeResult {
107 surviving_items: vec![],
108 findings: vec![],
109 rounds: 0,
110 converged: true,
111 round_stats: vec![],
112 });
113 }
114
115 tracing::info!(
116 n_items = items.len(),
117 max_rounds = config.max_rounds,
118 adversarial = config.adversarial,
119 "converge started"
120 );
121
122 let mut state = ConvergeState {
123 items,
124 findings: vec![],
125 round_stats: vec![],
126 converged: false,
127 };
128
129 for round in 1..=config.max_rounds {
130 let items_count = state.items.len();
131 tracing::info!(round, items_count, "converge round started");
132 let round_input = state.items.clone();
133
134 let (findings, _producer_stats) = generate_findings(
136 &round_input,
137 producer_prompt,
138 config.producers_per_item,
139 config.model.clone(),
140 scheduler,
141 run_ctx,
142 )
143 .await;
144
145 let current_findings = findings.clone();
146 state.findings.extend(current_findings);
147
148 if findings.is_empty() {
149 tracing::info!(round, "converge: no findings generated, ending");
150 break;
151 }
152
153 let (surviving_findings, vote_stats) = if config.adversarial {
155 tracing::debug!(
156 round,
157 n_findings = findings.len(),
158 adversaries = config.adversaries_per_finding,
159 "adversarial verification started"
160 );
161 verify_findings(
162 &findings,
163 adversary_prompt,
164 config.adversaries_per_finding,
165 config.vote_threshold,
166 config.model.clone(),
167 scheduler,
168 run_ctx,
169 )
170 .await
171 } else {
172 (findings.clone(), VoteStats::default())
173 };
174
175 let round_stats = RoundStats {
177 round,
178 items_input: items_count,
179 findings_generated: findings.len(),
180 findings_survived: surviving_findings.len(),
181 findings_refuted: findings.len() - surviving_findings.len(),
182 approval_rate: vote_stats.approval_rate,
183 };
184 tracing::info!(
185 round,
186 generated = findings.len(),
187 survived = surviving_findings.len(),
188 refuted = findings.len() - surviving_findings.len(),
189 approval_rate = vote_stats.approval_rate,
190 "converge round finished"
191 );
192 state.round_stats.push(round_stats);
193
194 if surviving_findings.is_empty() {
196 tracing::info!(round, "converge: all findings refuted, converged");
197 state.converged = true;
198 break;
199 }
200
201 if surviving_findings.len() == findings.len() && round > 1 {
203 tracing::info!(round, "converge: no findings refuted, full convergence");
204 state.converged = true;
205 break;
206 }
207
208 state.items = surviving_findings
210 .into_iter()
211 .map(|f| {
212 serde_json::json!({
213 "kind": f.kind,
214 "severity": format!("{:?}", f.severity).to_lowercase(),
215 "title": f.title,
216 "detail": f.detail,
217 "location": f.location,
218 "evidence": f.evidence,
219 "data": f.data
220 })
221 })
222 .collect();
223 }
224
225 tracing::info!(
226 rounds = state.round_stats.len(),
227 converged = state.converged,
228 surviving = state.items.len(),
229 total_findings = state.findings.len(),
230 "converge finished"
231 );
232 Ok(ConvergeResult {
233 surviving_items: state.items,
234 findings: state.findings,
235 rounds: state.round_stats.len() as u32,
236 converged: state.converged,
237 round_stats: state.round_stats,
238 })
239}
240
241async fn generate_findings(
243 items: &[serde_json::Value],
244 prompt_template: &str,
245 producers_per_item: u32,
246 model: Option<String>,
247 scheduler: &Arc<Scheduler>,
248 run_ctx: &RunContext,
249) -> (Vec<Finding>, ProducerStats) {
250 let mut tasks = Vec::new();
252 for item in items {
253 for _ in 0..producers_per_item {
254 let prompt =
255 prompt_template.replace("{item}", &serde_json::to_string(item).unwrap_or_default());
256 let agent_id = uuid::Uuid::now_v7();
257 let task = luft_core::contract::backend::AgentTask {
258 agent_id,
259 phase_id: 2, prompt,
261 model: model.clone(),
262 description: None,
263 role: Some("producer".to_string()),
264 name: None,
265 agent_seq: 0,
266 allowlist: None,
267 workdir: std::path::PathBuf::from("."),
268 mcp_endpoint: None,
269 timeout: None,
270 output_schema: None,
271 workdir_override: None,
272 };
273 tasks.push((task, None::<String>));
274 }
275 }
276
277 let results = scheduler.run_parallel(run_ctx.run_id, tasks).await;
279 let agents_run = results.len();
280 tracing::debug!(agents_run, "producer agents completed");
281 let mut all_findings = Vec::new();
282
283 for result in results.into_iter().flatten() {
284 all_findings.extend(result.findings);
285 }
286
287 let stats = ProducerStats {
288 items_processed: items.len(),
289 agents_run,
290 findings_generated: all_findings.len(),
291 };
292
293 (all_findings, stats)
294}
295
296async fn verify_findings(
298 findings: &[Finding],
299 prompt_template: &str,
300 adversaries_per_finding: u32,
301 vote_threshold: f32,
302 model: Option<String>,
303 scheduler: &Arc<Scheduler>,
304 run_ctx: &RunContext,
305) -> (Vec<Finding>, VoteStats) {
306 let mut votes: Vec<(Finding, usize)> = findings.iter().map(|f| (f.clone(), 0)).collect();
307 let total_votes_possible = findings.len() * adversaries_per_finding as usize;
308
309 for finding in findings {
310 let mut approval_count = 0usize;
311
312 for _ in 0..adversaries_per_finding {
313 let prompt = prompt_template.replace(
314 "{finding}",
315 &serde_json::to_string(finding).unwrap_or_default(),
316 );
317 let agent_id = uuid::Uuid::now_v7();
318
319 let task = luft_core::contract::backend::AgentTask {
320 agent_id,
321 phase_id: 2,
322 prompt,
323 model: model.clone(),
324 description: None,
325 role: Some("adversary".to_string()),
326 name: None,
327 agent_seq: 0,
328 allowlist: None,
329 workdir: std::path::PathBuf::from("."),
330 mcp_endpoint: None,
331 timeout: None,
332 output_schema: None,
333 workdir_override: None,
334 };
335
336 let result = scheduler.run_agent(run_ctx.run_id, task, None).await;
337
338 if let Ok(r) = result {
339 if r.status == AgentStatus::Ok {
340 tracing::trace!(%agent_id, "adversary approved finding");
341 approval_count += 1;
342 } else {
343 tracing::trace!(%agent_id, ?r.status, "adversary rejected finding");
344 }
345 }
346 }
347
348 if let Some((_, count)) = votes.iter_mut().find(|(f, _)| f.title == finding.title) {
349 *count = approval_count;
350 }
351 }
352
353 let approval_rate = if total_votes_possible > 0 {
354 votes.iter().map(|(_, c)| *c).sum::<usize>() as f32 / total_votes_possible as f32
355 } else {
356 0.0
357 };
358
359 let surviving: Vec<Finding> = votes
360 .into_iter()
361 .filter(|(_, count)| {
362 let threshold = (count * 100) / (adversaries_per_finding as usize).max(1);
363 threshold as f32 >= vote_threshold * 100.0
364 })
365 .map(|(f, _)| f)
366 .collect();
367
368 tracing::debug!(
369 surviving = surviving.len(),
370 refuted = findings.len() - surviving.len(),
371 approval_rate,
372 "adversarial voting completed"
373 );
374
375 let stats = VoteStats { approval_rate };
376
377 (surviving, stats)
378}
379
380#[derive(Debug, Default)]
381#[allow(dead_code)] struct ProducerStats {
383 items_processed: usize,
384 agents_run: usize,
385 findings_generated: usize,
386}
387
388#[derive(Debug, Default)]
389struct VoteStats {
390 approval_rate: f32,
391}
392
393fn extract_string(s: mlua::String) -> Option<String> {
399 s.to_str().ok().map(|s| s.to_string())
400}
401
402pub fn register_converge_sdk(lua: &Lua, cx: &SdkContext) -> mlua::Result<()> {
407 let globals = lua.globals();
408 let sched = cx.scheduler.clone();
409 let rc = cx.run_ctx.clone();
410 let handle = cx.handle.clone();
411 let events = cx.events();
412 let run_id = cx.run_id();
413 let phase_counter = cx.phase_counter.clone();
414 let span_counter = cx.span_counter.clone();
415
416 let converge_fn = lua.create_function(move |_lua, (items, options): (Table, Table)| {
417 let scheduler = sched.clone();
418 let run_ctx = rc.clone();
419 let handle = handle.clone();
420 let config = parse_converge_options(&options);
422 tracing::debug!(
423 adversarial = config.adversarial,
424 max_rounds = config.max_rounds,
425 "converge SDK invoked"
426 );
427 let phase_id = phase_counter.load(Ordering::Relaxed);
428 let span_id = span_counter.fetch_add(1, Ordering::Relaxed);
429 let max_rounds = config.max_rounds;
430
431 let producer_prompt = options
432 .get::<mlua::String>("producer_prompt")
433 .ok()
434 .and_then(extract_string)
435 .unwrap_or_else(|| {
436 "Analyze the following item and report any findings: {item}".to_string()
437 });
438
439 let adversary_prompt = options
440 .get::<mlua::String>("adversary_prompt")
441 .ok()
442 .and_then(extract_string)
443 .unwrap_or_else(|| {
444 "Review this finding and determine if it is valid or should be refuted: {finding}"
445 .to_string()
446 });
447
448 let items_vec: Vec<serde_json::Value> = items
450 .sequence_values()
451 .filter_map(|v: mlua::Result<Value>| v.ok())
452 .filter_map(|v| lua_value_to_json(&v).ok())
453 .collect();
454
455 let _ = events.send(AgentEvent::ConvergeStarted {
456 run_id,
457 phase_id,
458 span_id,
459 items: items_vec.len(),
460 max_rounds,
461 });
462 let t0 = std::time::Instant::now();
463
464 let result = handle.block_on(execute_convergence(
466 items_vec,
467 &producer_prompt,
468 &adversary_prompt,
469 config,
470 &scheduler,
471 &run_ctx,
472 ));
473 let elapsed_ms = t0.elapsed().as_millis() as u64;
474
475 match result {
476 Ok(ref res) => {
477 tracing::info!(
478 rounds = res.rounds,
479 converged = res.converged,
480 surviving = res.surviving_items.len(),
481 elapsed_ms,
482 "converge SDK completed"
483 );
484 }
485 Err(ref e) => {
486 tracing::error!(error = %e, elapsed_ms, "converge SDK failed");
487 }
488 }
489 match result {
490 Ok(res) => {
491 let _ = events.send(AgentEvent::ConvergeDone {
492 run_id,
493 phase_id,
494 span_id,
495 rounds: res.rounds,
496 converged: res.converged,
497 surviving: res.surviving_items.len(),
498 result: serde_json::json!({
499 "surviving": res.surviving_items,
500 "rounds": res.rounds,
501 "converged": res.converged,
502 "findings": res.findings,
503 }),
504 elapsed_ms,
505 error: None,
506 });
507 let result_table = _lua.create_table()?;
508 let surviving = _lua.create_table()?;
509 for (i, item) in res.surviving_items.iter().enumerate() {
510 let lua_val = json_to_lua_value(_lua, item.clone())?;
511 surviving.set(i + 1, lua_val)?;
512 }
513 result_table.set("surviving", surviving)?;
514 result_table.set("rounds", res.rounds)?;
515 result_table.set("converged", res.converged)?;
516
517 let findings_table = _lua.create_table()?;
518 for (i, finding) in res.findings.iter().enumerate() {
519 let ft = _lua.create_table()?;
520 ft.set("kind", finding.kind.as_str())?;
521 ft.set("severity", format!("{:?}", finding.severity).to_lowercase())?;
522 ft.set("title", finding.title.as_str())?;
523 ft.set("detail", finding.detail.as_str())?;
524 findings_table.set(i + 1, ft)?;
525 }
526 result_table.set("findings", findings_table)?;
527
528 Ok(result_table)
529 }
530 Err(e) => {
531 let _ = events.send(AgentEvent::ConvergeDone {
532 run_id,
533 phase_id,
534 span_id,
535 rounds: 0,
536 converged: false,
537 surviving: 0,
538 result: serde_json::Value::Null,
539 elapsed_ms,
540 error: Some(e.to_string()),
541 });
542 Err(mlua::Error::RuntimeError(format!("converge error: {}", e)))
543 }
544 }
545 })?;
546
547 globals.set("converge", converge_fn)?;
548 Ok(())
549}
550
551fn parse_converge_options(options: &Table) -> ConvergeConfig {
553 let mut config = ConvergeConfig::default();
554
555 if let Ok(adversarial) = options.get::<mlua::String>("adversarial") {
556 let s = extract_string(adversarial).unwrap_or_else(|| "true".to_string());
557 config.adversarial = s != "false";
558 }
559 if let Ok(threshold) = options.get::<mlua::Number>("vote_threshold") {
560 config.vote_threshold = threshold as f32;
561 }
562 if let Ok(max_rounds) = options.get::<mlua::Integer>("max_rounds") {
563 config.max_rounds = max_rounds as u32;
564 }
565 if let Ok(producers) = options.get::<mlua::Integer>("producers") {
566 config.producers_per_item = producers as u32;
567 }
568 if let Ok(adversaries) = options.get::<mlua::Integer>("adversaries") {
569 config.adversaries_per_finding = adversaries as u32;
570 }
571 if let Ok(model) = options.get::<mlua::String>("model") {
572 let s = extract_string(model).unwrap_or_default();
573 config.model = if s.is_empty() { None } else { Some(s) };
574 }
575
576 config
577}
578
579fn lua_value_to_json(value: &Value) -> Result<serde_json::Value, mlua::Error> {
581 match value {
582 Value::Nil => Ok(serde_json::Value::Null),
583 Value::Boolean(b) => Ok(serde_json::Value::Bool(*b)),
584 Value::Integer(i) => Ok(serde_json::Value::Number(serde_json::Number::from(*i))),
585 Value::Number(n) => {
586 if let Some(n) = serde_json::Number::from_f64(*n) {
587 Ok(serde_json::Value::Number(n))
588 } else {
589 Ok(serde_json::Value::Null)
590 }
591 }
592 Value::String(s) => {
593 let owned = s.clone();
594 let s = match owned.to_str() {
595 Ok(s) => s.to_string(),
596 Err(_) => return Ok(serde_json::Value::Null),
597 };
598 Ok(serde_json::Value::String(s))
599 }
600 Value::Table(t) => {
601 let len = t.len().unwrap_or(0);
603 if len > 0 {
604 let arr: Vec<serde_json::Value> = t
605 .sequence_values()
606 .filter_map(|v| v.ok())
607 .filter_map(|v| lua_value_to_json(&v).ok())
608 .collect();
609 if !arr.is_empty() {
610 return Ok(serde_json::Value::Array(arr));
611 }
612 }
613
614 let mut map = serde_json::Map::new();
616 for (k, v) in t.pairs::<Value, Value>().flatten() {
617 let key = match k {
618 Value::String(s) => s
619 .clone()
620 .to_str()
621 .map(|s| s.to_string())
622 .unwrap_or_default(),
623 Value::Integer(i) => i.to_string(),
624 _ => continue,
625 };
626 if let Ok(v) = lua_value_to_json(&v) {
627 map.insert(key, v);
628 }
629 }
630 Ok(serde_json::Value::Object(map))
631 }
632 _ => Ok(serde_json::Value::Null),
633 }
634}
635
636fn json_to_lua_value(lua: &Lua, json: serde_json::Value) -> Result<Value, mlua::Error> {
638 match json {
639 serde_json::Value::Null => Ok(Value::Nil),
640 serde_json::Value::Bool(b) => Ok(Value::Boolean(b)),
641 serde_json::Value::Number(n) => {
642 if let Some(i) = n.as_i64() {
643 Ok(Value::Integer(i))
644 } else if let Some(f) = n.as_f64() {
645 Ok(Value::Number(f))
646 } else {
647 Ok(Value::Nil)
648 }
649 }
650 serde_json::Value::String(s) => Ok(Value::String(lua.create_string(&s)?)),
651 serde_json::Value::Array(arr) => {
652 let t = lua.create_table()?;
653 for (i, v) in arr.into_iter().enumerate() {
654 t.set(i + 1, json_to_lua_value(lua, v)?)?;
655 }
656 Ok(Value::Table(t))
657 }
658 serde_json::Value::Object(map) => {
659 let t = lua.create_table()?;
660 for (k, v) in map {
661 t.set(k, json_to_lua_value(lua, v)?)?;
662 }
663 Ok(Value::Table(t))
664 }
665 }
666}
667
668#[cfg(test)]
669mod tests {
670 use super::*;
671 use crate::sdk::ReportSink;
672 use async_trait::async_trait;
673 use luft_core::contract::backend::{
674 AgentBackend, AgentCapabilities, AgentResult, BackendError, LogRef,
675 };
676 use luft_core::contract::finding::{Location, Severity};
677 use luft_core::contract::ids::TokenUsage;
678 use luft_core::scheduler::{BackendRegistry, RetryPolicy, SchedulerConfig};
679 use luft_core::{AgentTask, MockBackend, MockBehavior};
680 use std::sync::Arc;
681 use std::time::Duration;
682 use tokio::sync::broadcast;
683 use tokio_util::sync::CancellationToken;
684 use uuid::Uuid;
685
686 fn sample_finding(title: &str) -> Finding {
689 Finding {
690 kind: "test".into(),
691 severity: Severity::Info,
692 title: title.into(),
693 detail: "A detailed description".into(),
694 location: Some(Location {
695 file: "src/main.rs".into(),
696 line: Some(42),
697 }),
698 evidence: vec!["line 42: suspect code".into()],
699 data: serde_json::json!({"extra": "info"}),
700 }
701 }
702
703 fn default_finding() -> Finding {
704 sample_finding("Test Finding")
705 }
706
707 fn converge_scheduler(backend: Arc<dyn AgentBackend>) -> Arc<Scheduler> {
710 let config = SchedulerConfig {
711 max_concurrency: 4,
712 quota_per_run: 1000,
713 retry: RetryPolicy::default(),
714 };
715 let registry = BackendRegistry::new().with(backend);
716 Scheduler::new(config, registry, None)
717 }
718
719 fn test_run_ctx(scheduler: &Arc<Scheduler>) -> (RunContext, broadcast::Receiver<AgentEvent>) {
720 let run_id = Uuid::now_v7();
721 let rx = scheduler.init_run(run_id, 64);
722 let (tx, _rx2) = broadcast::channel(64);
723 let ctx = RunContext {
724 run_id,
725 cancel: CancellationToken::new(),
726 events: tx,
727 };
728 (ctx, rx)
729 }
730
731 #[test]
736 fn test_default_config() {
737 let c = ConvergeConfig::default();
738 assert!(c.adversarial);
739 assert!((c.vote_threshold - 0.7).abs() < f32::EPSILON);
740 assert_eq!(c.max_rounds, 3);
741 assert_eq!(c.producers_per_item, 1);
742 assert_eq!(c.adversaries_per_finding, 1);
743 assert!(c.model.is_none());
744 }
745
746 #[test]
747 fn test_converge_config_debug_clone_serialize() {
748 let c = ConvergeConfig::default();
749 let _ = format!("{:?}", c);
750 let _ = c.clone();
751 let json = serde_json::to_string(&c).unwrap();
752 let back: ConvergeConfig = serde_json::from_str(&json).unwrap();
753 assert_eq!(c.max_rounds, back.max_rounds);
754 }
755
756 #[test]
761 fn test_parse_options_defaults() {
762 let lua = Lua::new();
763 let t = lua.create_table().unwrap();
764 let cfg = parse_converge_options(&t);
765 assert!(cfg.adversarial);
766 assert!((cfg.vote_threshold - 0.7).abs() < f32::EPSILON);
767 assert_eq!(cfg.max_rounds, 3);
768 assert_eq!(cfg.producers_per_item, 1);
769 assert_eq!(cfg.adversaries_per_finding, 1);
770 assert!(cfg.model.is_none());
771 }
772
773 #[test]
774 fn test_parse_options_all_fields() {
775 let lua = Lua::new();
776 let t = lua.create_table().unwrap();
777 t.set("adversarial", "false").unwrap();
778 t.set("vote_threshold", 0.5).unwrap();
779 t.set("max_rounds", 5u32).unwrap();
780 t.set("producers", 2u32).unwrap();
781 t.set("adversaries", 3u32).unwrap();
782 t.set("model", "gpt-4").unwrap();
783 let cfg = parse_converge_options(&t);
784 assert!(!cfg.adversarial);
785 assert!((cfg.vote_threshold - 0.5).abs() < f32::EPSILON);
786 assert_eq!(cfg.max_rounds, 5);
787 assert_eq!(cfg.producers_per_item, 2);
788 assert_eq!(cfg.adversaries_per_finding, 3);
789 assert_eq!(cfg.model.as_deref(), Some("gpt-4"));
790 }
791
792 #[test]
793 fn test_parse_options_adversarial_true_string() {
794 let lua = Lua::new();
795 let t = lua.create_table().unwrap();
796 t.set("adversarial", "true").unwrap();
797 assert!(parse_converge_options(&t).adversarial);
798 }
799
800 #[test]
801 fn test_parse_options_empty_model() {
802 let lua = Lua::new();
803 let t = lua.create_table().unwrap();
804 t.set("model", "").unwrap();
805 assert!(parse_converge_options(&t).model.is_none());
806 }
807
808 #[test]
809 fn test_parse_options_model_some() {
810 let lua = Lua::new();
811 let t = lua.create_table().unwrap();
812 t.set("model", "claude").unwrap();
813 assert_eq!(parse_converge_options(&t).model.as_deref(), Some("claude"));
814 }
815
816 #[test]
821 fn test_extract_string_valid() {
822 let lua = Lua::new();
823 assert_eq!(
824 extract_string(lua.create_string("hello").unwrap()),
825 Some("hello".into())
826 );
827 }
828
829 #[test]
830 fn test_extract_string_empty() {
831 let lua = Lua::new();
832 assert_eq!(
833 extract_string(lua.create_string("").unwrap()),
834 Some("".into())
835 );
836 }
837
838 #[test]
839 fn test_extract_string_invalid_utf8() {
840 let lua = Lua::new();
841 let s = lua.create_string([0xFF, 0xFE, 0x00]).unwrap();
842 assert_eq!(extract_string(s), None);
843 }
844
845 #[test]
850 fn test_lua_value_to_json_nil() {
851 assert_eq!(
852 lua_value_to_json(&Value::Nil).unwrap(),
853 serde_json::Value::Null
854 );
855 }
856
857 #[test]
858 fn test_lua_value_to_json_boolean() {
859 assert_eq!(
860 lua_value_to_json(&Value::Boolean(true)).unwrap(),
861 serde_json::Value::Bool(true)
862 );
863 assert_eq!(
864 lua_value_to_json(&Value::Boolean(false)).unwrap(),
865 serde_json::Value::Bool(false)
866 );
867 }
868
869 #[test]
870 fn test_lua_value_to_json_integer() {
871 assert_eq!(
872 lua_value_to_json(&Value::Integer(42)).unwrap(),
873 serde_json::json!(42)
874 );
875 assert_eq!(
876 lua_value_to_json(&Value::Integer(-5)).unwrap(),
877 serde_json::json!(-5)
878 );
879 }
880
881 #[test]
882 fn test_lua_value_to_json_number() {
883 assert_eq!(
884 lua_value_to_json(&Value::Number(std::f64::consts::PI)).unwrap(),
885 serde_json::json!(std::f64::consts::PI)
886 );
887 }
888
889 #[test]
890 fn test_lua_value_to_json_number_nan() {
891 assert_eq!(
892 lua_value_to_json(&Value::Number(f64::NAN)).unwrap(),
893 serde_json::Value::Null
894 );
895 }
896
897 #[test]
898 fn test_lua_value_to_json_number_infinity() {
899 assert_eq!(
900 lua_value_to_json(&Value::Number(f64::INFINITY)).unwrap(),
901 serde_json::Value::Null
902 );
903 assert_eq!(
904 lua_value_to_json(&Value::Number(f64::NEG_INFINITY)).unwrap(),
905 serde_json::Value::Null
906 );
907 }
908
909 #[test]
910 fn test_lua_value_to_json_string_valid() {
911 let lua = Lua::new();
912 let s = lua.create_string("hello world").unwrap();
913 assert_eq!(
914 lua_value_to_json(&Value::String(s)).unwrap(),
915 serde_json::json!("hello world")
916 );
917 }
918
919 #[test]
920 fn test_lua_value_to_json_string_invalid_utf8() {
921 let lua = Lua::new();
922 let s = lua.create_string([0xFF, 0xFE]).unwrap();
923 assert_eq!(
924 lua_value_to_json(&Value::String(s)).unwrap(),
925 serde_json::Value::Null
926 );
927 }
928
929 #[test]
930 fn test_lua_value_to_json_table_as_array() {
931 let lua = Lua::new();
932 let t = lua.create_table().unwrap();
933 t.set(1, "a").unwrap();
934 t.set(2, "b").unwrap();
935 t.set(3, "c").unwrap();
936 assert_eq!(
937 lua_value_to_json(&Value::Table(t)).unwrap(),
938 serde_json::json!(["a", "b", "c"])
939 );
940 }
941
942 #[test]
943 fn test_lua_value_to_json_table_as_object() {
944 let lua = Lua::new();
945 let t = lua.create_table().unwrap();
946 t.set("name", "test").unwrap();
947 t.set("count", 42).unwrap();
948 let val = lua_value_to_json(&Value::Table(t)).unwrap();
949 let obj = val.as_object().unwrap();
950 assert_eq!(obj["name"], "test");
951 assert_eq!(obj["count"], 42);
952 }
953
954 #[test]
955 fn test_lua_value_to_json_table_empty() {
956 let lua = Lua::new();
957 let t = lua.create_table().unwrap();
958 assert_eq!(
959 lua_value_to_json(&Value::Table(t)).unwrap(),
960 serde_json::json!({})
961 );
962 }
963
964 #[test]
965 fn test_lua_value_to_json_table_nested() {
966 let lua = Lua::new();
967 let inner = lua.create_table().unwrap();
968 inner.set("key", "val").unwrap();
969 let outer = lua.create_table().unwrap();
970 outer.set("nested", inner).unwrap();
971 let val = lua_value_to_json(&Value::Table(outer)).unwrap();
972 assert_eq!(val, serde_json::json!({"nested": {"key": "val"}}));
973 }
974
975 #[test]
976 fn test_lua_value_to_json_table_integer_key_in_object() {
977 let lua = Lua::new();
978 let t = lua.create_table().unwrap();
979 t.set(42, "answer").unwrap();
980 let val = lua_value_to_json(&Value::Table(t)).unwrap();
981 assert_eq!(val, serde_json::json!({"42": "answer"}));
982 }
983
984 #[test]
985 fn test_lua_value_to_json_function_catch_all() {
986 let lua = Lua::new();
987 let f = lua.create_function(|_, ()| Ok(())).unwrap();
988 assert_eq!(
989 lua_value_to_json(&Value::Function(f)).unwrap(),
990 serde_json::Value::Null
991 );
992 }
993
994 #[test]
995 fn test_lua_value_to_json_userdata_catch_all() {
996 let lua = Lua::new();
997 let thread = lua
999 .create_thread(lua.load("return 1").into_function().unwrap())
1000 .unwrap();
1001 let val = Value::Thread(thread);
1002 assert_eq!(lua_value_to_json(&val).unwrap(), serde_json::Value::Null);
1003 }
1004
1005 #[test]
1010 fn test_json_to_lua_value_null() {
1011 let lua = Lua::new();
1012 assert!(matches!(
1013 json_to_lua_value(&lua, serde_json::Value::Null).unwrap(),
1014 Value::Nil
1015 ));
1016 }
1017
1018 #[test]
1019 fn test_json_to_lua_value_bool() {
1020 let lua = Lua::new();
1021 assert!(matches!(
1022 json_to_lua_value(&lua, serde_json::Value::Bool(true)).unwrap(),
1023 Value::Boolean(true)
1024 ));
1025 }
1026
1027 #[test]
1028 fn test_json_to_lua_value_integer() {
1029 let lua = Lua::new();
1030 assert!(matches!(
1031 json_to_lua_value(&lua, serde_json::json!(42)).unwrap(),
1032 Value::Integer(42)
1033 ));
1034 }
1035
1036 #[test]
1037 fn test_json_to_lua_value_float() {
1038 let lua = Lua::new();
1039 assert!(matches!(
1040 json_to_lua_value(&lua, serde_json::json!(std::f64::consts::PI)).unwrap(),
1041 Value::Number(n) if (n - std::f64::consts::PI).abs() < f64::EPSILON
1042 ));
1043 }
1044
1045 #[test]
1046 fn test_json_to_lua_value_large_number() {
1047 let lua = Lua::new();
1048 let val = json_to_lua_value(&lua, serde_json::json!(1e200)).unwrap();
1050 assert!(matches!(val, Value::Number(_)));
1051 }
1052
1053 #[test]
1054 fn test_json_to_lua_value_string() {
1055 let lua = Lua::new();
1056 assert!(matches!(
1057 json_to_lua_value(&lua, serde_json::Value::String("hi".into())).unwrap(),
1058 Value::String(_)
1059 ));
1060 }
1061
1062 #[test]
1063 fn test_json_to_lua_value_array() {
1064 let lua = Lua::new();
1065 let val = json_to_lua_value(&lua, serde_json::json!([1, 2, 3])).unwrap();
1066 assert!(matches!(val, Value::Table(_)));
1067 }
1068
1069 #[test]
1070 fn test_json_to_lua_value_empty_array() {
1071 let lua = Lua::new();
1072 let val = json_to_lua_value(&lua, serde_json::json!([])).unwrap();
1073 assert!(matches!(val, Value::Table(_)));
1074 }
1075
1076 #[test]
1077 fn test_json_to_lua_value_object() {
1078 let lua = Lua::new();
1079 let val = json_to_lua_value(&lua, serde_json::json!({"a": 1})).unwrap();
1080 assert!(matches!(val, Value::Table(_)));
1081 }
1082
1083 #[test]
1084 fn test_json_to_lua_value_empty_object() {
1085 let lua = Lua::new();
1086 let val = json_to_lua_value(&lua, serde_json::json!({})).unwrap();
1087 assert!(matches!(val, Value::Table(_)));
1088 }
1089
1090 #[test]
1091 fn test_json_to_lua_value_nested() {
1092 let lua = Lua::new();
1093 let val = json_to_lua_value(&lua, serde_json::json!({"a": {"b": [1, 2, 3]}})).unwrap();
1094 assert!(matches!(val, Value::Table(_)));
1095 }
1096
1097 #[test]
1102 fn test_converge_result_construction() {
1103 let r = ConvergeResult {
1104 surviving_items: vec![serde_json::json!({"k": "v"})],
1105 findings: vec![default_finding()],
1106 rounds: 2,
1107 converged: true,
1108 round_stats: vec![RoundStats {
1109 round: 1,
1110 items_input: 3,
1111 findings_generated: 5,
1112 findings_survived: 4,
1113 findings_refuted: 1,
1114 approval_rate: 0.8,
1115 }],
1116 };
1117 assert_eq!(r.rounds, 2);
1118 assert!(r.converged);
1119 assert_eq!(r.surviving_items.len(), 1);
1120 assert_eq!(r.findings.len(), 1);
1121 assert_eq!(r.round_stats.len(), 1);
1122 }
1123
1124 #[test]
1125 fn test_converge_result_serialize_roundtrip() {
1126 let r = ConvergeResult {
1127 surviving_items: vec![],
1128 findings: vec![],
1129 rounds: 0,
1130 converged: true,
1131 round_stats: vec![],
1132 };
1133 let json = serde_json::to_string(&r).unwrap();
1134 let back: ConvergeResult = serde_json::from_str(&json).unwrap();
1135 assert_eq!(back.rounds, 0);
1136 assert!(back.converged);
1137 }
1138
1139 #[test]
1140 fn test_converge_result_debug_clone() {
1141 let r = ConvergeResult {
1142 surviving_items: vec![],
1143 findings: vec![],
1144 rounds: 1,
1145 converged: false,
1146 round_stats: vec![],
1147 };
1148 let _ = format!("{:?}", r);
1149 let _ = r.clone();
1150 }
1151
1152 #[test]
1153 fn test_round_stats_construction() {
1154 let s = RoundStats {
1155 round: 1,
1156 items_input: 10,
1157 findings_generated: 20,
1158 findings_survived: 15,
1159 findings_refuted: 5,
1160 approval_rate: 0.75,
1161 };
1162 assert_eq!(s.round, 1);
1163 assert_eq!(s.items_input, 10);
1164 assert_eq!(s.findings_generated, 20);
1165 assert_eq!(s.findings_survived, 15);
1166 assert_eq!(s.findings_refuted, 5);
1167 assert!((s.approval_rate - 0.75).abs() < f32::EPSILON);
1168 }
1169
1170 #[test]
1171 fn test_round_stats_serialize_roundtrip() {
1172 let s = RoundStats {
1173 round: 2,
1174 items_input: 5,
1175 findings_generated: 8,
1176 findings_survived: 3,
1177 findings_refuted: 5,
1178 approval_rate: 0.4,
1179 };
1180 let json = serde_json::to_string(&s).unwrap();
1181 let back: RoundStats = serde_json::from_str(&json).unwrap();
1182 assert_eq!(back.round, 2);
1183 assert!((back.approval_rate - 0.4).abs() < f32::EPSILON);
1184 }
1185
1186 #[tokio::test]
1191 async fn test_execute_convergence_empty_items() {
1192 let scheduler = converge_scheduler(Arc::new(NoOpBackend));
1193 let (ctx, _rx) = test_run_ctx(&scheduler);
1194
1195 let result = execute_convergence(
1196 vec![],
1197 "producer: {item}",
1198 "adversary: {finding}",
1199 ConvergeConfig::default(),
1200 &scheduler,
1201 &ctx,
1202 )
1203 .await
1204 .unwrap();
1205
1206 assert!(result.surviving_items.is_empty());
1207 assert!(result.findings.is_empty());
1208 assert_eq!(result.rounds, 0);
1209 assert!(result.converged);
1210 assert!(result.round_stats.is_empty());
1211 }
1212
1213 #[tokio::test]
1218 async fn test_execute_convergence_no_findings() {
1219 let mock = Arc::new(MockBackend::new(
1220 "mock",
1221 vec![MockBehavior::Success {
1222 output: serde_json::Value::Null,
1223 tokens: TokenUsage::default(),
1224 delay: Duration::from_millis(1),
1225 }],
1226 ));
1227 let scheduler = converge_scheduler(mock);
1228 let (ctx, _rx) = test_run_ctx(&scheduler);
1229
1230 let result = execute_convergence(
1231 vec![serde_json::json!({"key": "val"})],
1232 "producer: {item}",
1233 "adversary: {finding}",
1234 ConvergeConfig::default(),
1235 &scheduler,
1236 &ctx,
1237 )
1238 .await
1239 .unwrap();
1240
1241 assert_eq!(result.rounds, 0);
1244 assert!(!result.converged);
1245 assert!(result.findings.is_empty());
1246 assert_eq!(
1247 result.surviving_items,
1248 vec![serde_json::json!({"key": "val"})]
1249 );
1250 assert!(result.round_stats.is_empty());
1251 }
1252
1253 #[tokio::test]
1258 async fn test_execute_convergence_findings_survive() {
1259 let findings = vec![default_finding()];
1260 let backend = Arc::new(FindingsAlwaysBackend { findings });
1261 let scheduler = converge_scheduler(backend);
1262 let (ctx, _rx) = test_run_ctx(&scheduler);
1263
1264 let result = execute_convergence(
1265 vec![serde_json::json!({"input": "item"})],
1266 "producer: {item}",
1267 "adversary: {finding}",
1268 ConvergeConfig::default(),
1269 &scheduler,
1270 &ctx,
1271 )
1272 .await
1273 .unwrap();
1274
1275 assert_eq!(result.rounds, 2);
1278 assert!(result.converged);
1279 assert!(!result.findings.is_empty());
1280 assert_eq!(result.round_stats.len(), 2);
1281 for rs in &result.round_stats {
1283 assert_eq!(rs.findings_generated, 1);
1284 assert_eq!(rs.findings_survived, 1);
1285 assert_eq!(rs.findings_refuted, 0);
1286 }
1287 }
1288
1289 #[tokio::test]
1294 async fn test_execute_convergence_non_adversarial() {
1295 let findings = vec![default_finding()];
1296 let backend = Arc::new(FindingsAlwaysBackend { findings });
1297 let scheduler = converge_scheduler(backend);
1298 let (ctx, _rx) = test_run_ctx(&scheduler);
1299
1300 let config = ConvergeConfig {
1301 adversarial: false,
1302 ..ConvergeConfig::default()
1303 };
1304
1305 let result = execute_convergence(
1306 vec![serde_json::json!({"x": 1})],
1307 "producer: {item}",
1308 "adversary: {finding}",
1309 config,
1310 &scheduler,
1311 &ctx,
1312 )
1313 .await
1314 .unwrap();
1315
1316 assert_eq!(result.rounds, 2);
1318 assert!(result.converged);
1319 assert!(!result.findings.is_empty());
1320 }
1321
1322 #[tokio::test]
1327 async fn test_execute_convergence_all_refuted() {
1328 let findings = vec![default_finding()];
1329 let backend = Arc::new(FindingsRefutingBackend { findings });
1330 let scheduler = converge_scheduler(backend);
1331 let (ctx, _rx) = test_run_ctx(&scheduler);
1332
1333 let result = execute_convergence(
1334 vec![serde_json::json!({"x": 1})],
1335 "producer: {item}",
1336 "adversary: {finding}",
1337 ConvergeConfig::default(),
1338 &scheduler,
1339 &ctx,
1340 )
1341 .await
1342 .unwrap();
1343
1344 assert_eq!(result.rounds, 1);
1346 assert!(result.converged);
1347 assert_eq!(result.round_stats.len(), 1);
1348 assert_eq!(result.round_stats[0].findings_generated, 1);
1349 assert_eq!(result.round_stats[0].findings_survived, 0);
1350 assert_eq!(result.round_stats[0].findings_refuted, 1);
1351 }
1352
1353 #[tokio::test]
1358 async fn test_execute_convergence_no_adversarial_findings_generated() {
1359 let findings = vec![default_finding()];
1360 let backend = Arc::new(FindingsAlwaysBackend { findings });
1361 let scheduler = converge_scheduler(backend);
1362 let (ctx, _rx) = test_run_ctx(&scheduler);
1363
1364 let config = ConvergeConfig {
1365 adversarial: false,
1366 max_rounds: 1,
1367 ..ConvergeConfig::default()
1368 };
1369
1370 let result = execute_convergence(
1371 vec![serde_json::json!({"x": 1})],
1372 "producer: {item}",
1373 "adversary: {finding}",
1374 config,
1375 &scheduler,
1376 &ctx,
1377 )
1378 .await
1379 .unwrap();
1380
1381 assert_eq!(result.rounds, 1);
1382 assert!(!result.converged);
1384 assert_eq!(result.findings.len(), 1);
1385 }
1386
1387 #[test]
1392 fn test_register_converge_sdk_empty_items() {
1393 let rt = tokio::runtime::Runtime::new().unwrap();
1394 let lua = Lua::new();
1395 let scheduler = converge_scheduler(Arc::new(NoOpBackend));
1396 let run_id = Uuid::now_v7();
1397 let (tx, _rx2) = broadcast::channel(64);
1398 let run_ctx = RunContext {
1399 run_id,
1400 cancel: CancellationToken::new(),
1401 events: tx,
1402 };
1403 let report_sink: ReportSink = Arc::new(std::sync::Mutex::new(None));
1404 let handle = rt.handle().clone();
1405 let cx = SdkContext::new(run_ctx, scheduler, report_sink, None, handle);
1406
1407 register_converge_sdk(&lua, &cx).unwrap();
1408
1409 let globals = lua.globals();
1410 let converge: mlua::Function = globals.get("converge").unwrap();
1411
1412 let items = lua.create_table().unwrap();
1413 let options = lua.create_table().unwrap();
1414 let result: mlua::Table = converge.call((items, options)).unwrap();
1415
1416 let converged: bool = result.get("converged").unwrap();
1417 assert!(converged);
1418 let rounds: u32 = result.get("rounds").unwrap();
1419 assert_eq!(rounds, 0);
1420 let surviving: mlua::Table = result.get("surviving").unwrap();
1421 assert_eq!(surviving.len().unwrap(), 0);
1422 let findings: mlua::Table = result.get("findings").unwrap();
1423 assert_eq!(findings.len().unwrap(), 0);
1424 }
1425
1426 #[test]
1427 fn test_register_converge_sdk_with_items() {
1428 let rt = tokio::runtime::Runtime::new().unwrap();
1429 let lua = Lua::new();
1430 let backend = Arc::new(FindingsAlwaysBackend {
1431 findings: vec![default_finding()],
1432 });
1433 let scheduler = converge_scheduler(backend);
1434 let run_id = Uuid::now_v7();
1435 let _rx = scheduler.init_run(run_id, 64);
1436 let (tx, _rx2) = broadcast::channel(64);
1437 let run_ctx = RunContext {
1438 run_id,
1439 cancel: CancellationToken::new(),
1440 events: tx,
1441 };
1442 let report_sink: ReportSink = Arc::new(std::sync::Mutex::new(None));
1443 let handle = rt.handle().clone();
1444 let cx = SdkContext::new(run_ctx, scheduler, report_sink, None, handle);
1445
1446 register_converge_sdk(&lua, &cx).unwrap();
1447
1448 let globals = lua.globals();
1449 let converge: mlua::Function = globals.get("converge").unwrap();
1450
1451 let items = lua.create_table().unwrap();
1452 items.set(1, "test item").unwrap();
1453 let options = lua.create_table().unwrap();
1454 let result: mlua::Table = converge.call((items, options)).unwrap();
1455
1456 let converged: bool = result.get("converged").unwrap();
1457 assert!(converged);
1458 let rounds: u32 = result.get("rounds").unwrap();
1459 assert!(rounds > 0);
1460 let findings: mlua::Table = result.get("findings").unwrap();
1461 assert!(findings.len().unwrap() > 0);
1462 }
1463
1464 #[test]
1469 fn test_producer_stats() {
1470 let s = ProducerStats {
1471 items_processed: 5,
1472 agents_run: 10,
1473 findings_generated: 20,
1474 };
1475 assert_eq!(s.items_processed, 5);
1476 assert_eq!(s.agents_run, 10);
1477 assert_eq!(s.findings_generated, 20);
1478 }
1479
1480 #[test]
1481 fn test_producer_stats_default() {
1482 let s = ProducerStats::default();
1483 assert_eq!(s.items_processed, 0);
1484 assert_eq!(s.agents_run, 0);
1485 assert_eq!(s.findings_generated, 0);
1486 }
1487
1488 #[test]
1489 fn test_vote_stats() {
1490 let s = VoteStats { approval_rate: 0.5 };
1491 assert!((s.approval_rate - 0.5).abs() < f32::EPSILON);
1492 }
1493
1494 #[test]
1495 fn test_vote_stats_default() {
1496 let s = VoteStats::default();
1497 assert!((s.approval_rate - 0.0).abs() < f32::EPSILON);
1498 }
1499
1500 struct NoOpBackend;
1507
1508 #[async_trait]
1509 impl AgentBackend for NoOpBackend {
1510 fn id(&self) -> &'static str {
1511 "noop"
1512 }
1513 fn capabilities(&self) -> AgentCapabilities {
1514 AgentCapabilities::default()
1515 }
1516 fn as_any(&self) -> &dyn std::any::Any {
1517 self
1518 }
1519 async fn run(
1520 &self,
1521 task: AgentTask,
1522 _ctx: RunContext,
1523 ) -> Result<AgentResult, BackendError> {
1524 Ok(AgentResult {
1525 agent_id: task.agent_id,
1526 status: AgentStatus::Ok,
1527 output: serde_json::Value::Null,
1528 findings: vec![],
1529 tokens_used: TokenUsage::default(),
1530 artifacts: vec![],
1531 logs: LogRef::default(),
1532 })
1533 }
1534 }
1535
1536 struct FindingsAlwaysBackend {
1539 findings: Vec<Finding>,
1540 }
1541
1542 #[async_trait]
1543 impl AgentBackend for FindingsAlwaysBackend {
1544 fn id(&self) -> &'static str {
1545 "findings-always"
1546 }
1547 fn capabilities(&self) -> AgentCapabilities {
1548 AgentCapabilities::default()
1549 }
1550 fn as_any(&self) -> &dyn std::any::Any {
1551 self
1552 }
1553 async fn run(
1554 &self,
1555 task: AgentTask,
1556 _ctx: RunContext,
1557 ) -> Result<AgentResult, BackendError> {
1558 Ok(AgentResult {
1559 agent_id: task.agent_id,
1560 status: AgentStatus::Ok,
1561 output: serde_json::Value::Null,
1562 findings: self.findings.clone(),
1563 tokens_used: TokenUsage::default(),
1564 artifacts: vec![],
1565 logs: LogRef::default(),
1566 })
1567 }
1568 }
1569
1570 struct FindingsRefutingBackend {
1573 findings: Vec<Finding>,
1574 }
1575
1576 #[async_trait]
1577 impl AgentBackend for FindingsRefutingBackend {
1578 fn id(&self) -> &'static str {
1579 "findings-refute"
1580 }
1581 fn capabilities(&self) -> AgentCapabilities {
1582 AgentCapabilities::default()
1583 }
1584 fn as_any(&self) -> &dyn std::any::Any {
1585 self
1586 }
1587 async fn run(
1588 &self,
1589 task: AgentTask,
1590 _ctx: RunContext,
1591 ) -> Result<AgentResult, BackendError> {
1592 Ok(AgentResult {
1593 agent_id: task.agent_id,
1594 status: AgentStatus::Error,
1595 output: serde_json::Value::Null,
1596 findings: self.findings.clone(),
1597 tokens_used: TokenUsage::default(),
1598 artifacts: vec![],
1599 logs: LogRef::default(),
1600 })
1601 }
1602 }
1603}