diff --git a/core/codegen/src/cli_gen.rs b/core/codegen/src/cli_gen.rs index d93019cea41..144ea23133c 100644 --- a/core/codegen/src/cli_gen.rs +++ b/core/codegen/src/cli_gen.rs @@ -57,6 +57,12 @@ pub struct ToolMeta { /// and `--mode=ensure` (dev: write if changed, default). In verify mode, /// the CLI forces dry-run execution so content_upsert nodes check but don't write. pub enable_mode: bool, + /// Available profile names for `--profile` enum flag (C20/RT59). + /// + /// When non-empty, the generated CLI accepts `--profile ` to select + /// which interface bindings are active. Profile selection determines runtime + /// behavior for services bound via `profile { bind Interface { impl: ... } }`. + pub available_profiles: Vec, } /// An entrypoint that becomes a CLI flag. @@ -406,9 +412,24 @@ fn build_cli_imports(tool: &ToolMeta, custom_import: Option<&str>, step_mode: bo // ============================================================================ /// Generate the graph builder call expression. -fn generate_graph_builder_call(tool: &ToolMeta) -> String { +/// +/// When `has_profiles` is true, appends `selected_profile.as_deref()` to the +/// graph builder arguments so the profile selection flows through to the DAG. +fn generate_graph_builder_call(tool: &ToolMeta, has_profiles: bool) -> String { let f = &tool.graph_builder_call; - let args = &tool.graph_builder_args; + let base_args = &tool.graph_builder_args; + + // Build final args list, optionally including profile + let args = if has_profiles { + if base_args.is_empty() { + "selected_profile.as_deref()".to_string() + } else { + format!("{}, selected_profile.as_deref()", base_args) + } + } else { + base_args.to_string() + }; + if tool.returns_result { let call = if args.is_empty() { format!("{}()", f) @@ -433,9 +454,11 @@ fn generate_graph_builder_call(tool: &ToolMeta) -> String { /// by name) still compiles. /// /// When `enable_mode` is true, adds a `--mode` string parameter to the schema. +/// When `available_profiles` is non-empty, adds a `--profile` string parameter. fn generate_arg_parsing_with_mode( entrypoints: &[CliEntrypoint], enable_mode: bool, + available_profiles: &[String], ) -> String { let mut code = String::new(); @@ -446,6 +469,11 @@ fn generate_arg_parsing_with_mode( " gunbc_cli::CliParam::new(\"mode\", gunbc_cli::ParamType::Str),\n", ); } + if !available_profiles.is_empty() { + code.push_str( + " gunbc_cli::CliParam::new(\"profile\", gunbc_cli::ParamType::Str),\n", + ); + } for ep in entrypoints { let type_expr = match ep.type_id { ParamType::Str => "gunbc_cli::ParamType::Str", @@ -722,6 +750,45 @@ fn generate_mode_block(tool: &ToolMeta) -> String { code } +/// Generate the `--profile` handling block (C20/RT59). +/// +/// When `available_profiles` is non-empty, generates code that: +/// 1. Extracts the `--profile` value from parsed CLI args +/// 2. Validates it against the available profile names +/// 3. Exits with an error if the profile is invalid +/// +/// The profile value is used when building the graph with `build_dsl_graph_with_profile`. +fn generate_profile_block(tool: &ToolMeta) -> String { + if tool.available_profiles.is_empty() { + return String::new(); + } + + let profiles_list = tool + .available_profiles + .iter() + .map(|p| format!("\"{}\"", p)) + .collect::>() + .join(", "); + + let mut code = String::new(); + code.push_str("// --profile flag: select interface bindings (C20/RT59)\n"); + code.push_str(&format!( + "let valid_profiles: &[&str] = &[{}];\n", + profiles_list + )); + code.push_str("let selected_profile: Option = match cli_inputs.get(\"profile\") {\n"); + code.push_str(" Some(Value::Str(p)) => {\n"); + code.push_str(" if !valid_profiles.contains(&p.as_str()) {\n"); + code.push_str(" eprintln!(\"error: invalid profile '{}'. Valid profiles: {:?}\", p, valid_profiles);\n"); + code.push_str(" process::exit(1);\n"); + code.push_str(" }\n"); + code.push_str(" Some(p.clone())\n"); + code.push_str(" }\n"); + code.push_str(" _ => None, // default: no profile selected (stub interfaces)\n"); + code.push_str("};\n\n"); + code +} + /// Generate the dry-run mode block. fn generate_dry_run_block(tool: &ToolMeta) -> String { let mock_setup = generate_mock_setup(&tool.mock_spec_call); @@ -771,13 +838,15 @@ fn build_cli_source_file( /// Build the `main()` function for standard mode. fn build_main_fn(tool: &ToolMeta, entrypoints: &[CliEntrypoint]) -> FnDef { - let arg_parsing = generate_arg_parsing_with_mode(entrypoints, tool.enable_mode); - let graph_builder_call = generate_graph_builder_call(tool); + let has_profiles = !tool.available_profiles.is_empty(); + let arg_parsing = generate_arg_parsing_with_mode(entrypoints, tool.enable_mode, &tool.available_profiles); + let graph_builder_call = generate_graph_builder_call(tool, has_profiles); let input_mocks = generate_input_mocks(entrypoints); let dry_run_block = generate_dry_run_block(tool); let body_lines_expr = generate_preamble_body_lines(entrypoints); let success_port_arg = generate_success_port_arg(tool); let mode_block = generate_mode_block(tool); + let profile_block = generate_profile_block(tool); let body_code = format!( "let args: Vec = env::args().collect();\n\ @@ -785,6 +854,7 @@ fn build_main_fn(tool: &ToolMeta, entrypoints: &[CliEntrypoint]) -> FnDef { // Parse arguments\n\ {arg_parsing}\n\ {mode_block}\ + {profile_block}\ // Build the graph and compose with freshness checks\n\ let dag = {graph_builder_call};\n\ let steps = check_and_plan_freshness();\n\ @@ -804,6 +874,7 @@ fn build_main_fn(tool: &ToolMeta, entrypoints: &[CliEntrypoint]) -> FnDef { execute_and_display(&dag, mode, animated, {success_port_arg}, Some(&input_mocks));", arg_parsing = arg_parsing, mode_block = mode_block, + profile_block = profile_block, graph_builder_call = graph_builder_call, input_mocks = input_mocks, dry_run_block = dry_run_block, @@ -834,6 +905,16 @@ fn build_help_fn(tool: &ToolMeta, entrypoints: &[CliEntrypoint]) -> FnDef { "" }; + let profile_help = if tool.available_profiles.is_empty() { + String::new() + } else { + let profiles = tool.available_profiles.join(", "); + format!( + "println!(\" --profile NAME Select profile ({})\");\n ", + profiles + ) + }; + let body_code = format!( "println!(\"{tool_name} - {description}\");\n\ println!();\n\ @@ -844,6 +925,7 @@ fn build_help_fn(tool: &ToolMeta, entrypoints: &[CliEntrypoint]) -> FnDef { {help_options}\ println!(\" -n, --dry-run Don't perform actual I/O\");\n\ {mode_help}\ + {profile_help}\ println!(\" --print-inputs json Print parsed inputs as JSON and exit\");\n\ println!(\" -h, --help Print this help\");\n\ println!();\n\ @@ -852,6 +934,7 @@ fn build_help_fn(tool: &ToolMeta, entrypoints: &[CliEntrypoint]) -> FnDef { description = tool.description, help_options = help_options, mode_help = mode_help, + profile_help = profile_help, ); FnDef { @@ -963,24 +1046,37 @@ fn build_subcmd_run_fn( tool: &ToolMeta, subcmd: &crate::registry::SubcommandDef, ) -> FnDef { - let arg_parsing = generate_arg_parsing_with_mode(&subcmd.entrypoints, tool.enable_mode); + let has_profiles = !tool.available_profiles.is_empty(); + let arg_parsing = generate_arg_parsing_with_mode(&subcmd.entrypoints, tool.enable_mode, &tool.available_profiles); let input_mocks = generate_input_mocks(&subcmd.entrypoints); let body_lines_expr = generate_preamble_body_lines(&subcmd.entrypoints); + // Build args, optionally including profile + let base_args = &subcmd.graph_builder_args; + let args = if has_profiles { + if base_args.is_empty() { + "selected_profile.as_deref()".to_string() + } else { + format!("{}, selected_profile.as_deref()", base_args) + } + } else { + base_args.to_string() + }; + let graph_builder_call = if subcmd.returns_result { - let call = if subcmd.graph_builder_args.is_empty() { + let call = if args.is_empty() { format!("{}()", subcmd.graph_builder_call) } else { - format!("{}({})", subcmd.graph_builder_call, subcmd.graph_builder_args) + format!("{}({})", subcmd.graph_builder_call, args) }; format!( "match {} {{\n Ok(d) => d,\n Err(e) => {{\n eprintln!(\"Error building graph: {{}}\", e);\n process::exit(1);\n }}\n}}", call ) - } else if subcmd.graph_builder_args.is_empty() { + } else if args.is_empty() { format!("{}()", subcmd.graph_builder_call) } else { - format!("{}({})", subcmd.graph_builder_call, subcmd.graph_builder_args) + format!("{}({})", subcmd.graph_builder_call, args) }; let mock_setup = match &subcmd.mock_spec_call { @@ -1006,6 +1102,8 @@ fn build_subcmd_run_fn( String::new() }; + let profile_block = generate_profile_block(tool); + let body_code = format!( "// Reconstruct args with program name for parser compatibility\n\ let mut args: Vec = Vec::new();\n\ @@ -1015,6 +1113,7 @@ fn build_subcmd_run_fn( \n\ {arg_parsing}\n\ {mode_block}\ + {profile_block}\ // Build the graph and compose with freshness checks\n\ let dag = {graph_builder_call};\n\ let steps = check_and_plan_freshness();\n\ @@ -1034,6 +1133,7 @@ fn build_subcmd_run_fn( subcmd_name = subcmd.name, arg_parsing = arg_parsing, mode_block = mode_block, + profile_block = profile_block, graph_builder_call = graph_builder_call, input_mocks = input_mocks, dry_run_block = dry_run_block, @@ -1183,8 +1283,9 @@ match parsed.subcommand {\n\ /// Build the `run_full_dag()` function for step mode. fn build_run_full_dag_fn(tool: &ToolMeta, entrypoints: &[CliEntrypoint]) -> FnDef { - let arg_parsing = generate_arg_parsing_with_mode(entrypoints, false); - let graph_builder_call = generate_graph_builder_call(tool); + // Step mode doesn't need profile support - it's for CI step execution + let arg_parsing = generate_arg_parsing_with_mode(entrypoints, false, &[]); + let graph_builder_call = generate_graph_builder_call(tool, false); let input_mocks = generate_input_mocks(entrypoints); let dry_run_block = generate_dry_run_block(tool); let body_lines_expr = generate_preamble_body_lines(entrypoints); @@ -1236,7 +1337,8 @@ fn build_run_full_dag_fn(tool: &ToolMeta, entrypoints: &[CliEntrypoint]) -> FnDe /// Build the `run_single_step()` function for step mode. fn build_run_single_step_fn(tool: &ToolMeta) -> FnDef { - let graph_builder_call = generate_graph_builder_call(tool); + // Step mode doesn't need profile support - it's for CI step execution + let graph_builder_call = generate_graph_builder_call(tool, false); let dry_run_block = generate_dry_run_block(tool); let success_port_or_empty = tool.success_port.as_deref().unwrap_or(""); @@ -1323,7 +1425,8 @@ fn build_run_single_step_fn(tool: &ToolMeta) -> FnDef { /// Build the `list_dag_steps()` function for step mode. fn build_list_dag_steps_fn(tool: &ToolMeta) -> FnDef { - let graph_builder_call = generate_graph_builder_call(tool); + // Step mode doesn't need profile support - it's for CI step execution + let graph_builder_call = generate_graph_builder_call(tool, false); let body_code = format!( "let dag = {graph_builder_call};\n\ @@ -1546,6 +1649,7 @@ mod tests { enable_step_mode: false, mock_spec_call: Some("some_crate::graph_mock::mock_spec()".into()), enable_mode: false, + available_profiles: vec![], }; let entrypoints = vec![CliEntrypoint::new("repo_path", ParamType::Str) @@ -1577,6 +1681,7 @@ mod tests { enable_step_mode: false, mock_spec_call: Some("mock_spec()".into()), enable_mode: false, + available_profiles: vec![], }; let entrypoints = vec![]; @@ -1604,6 +1709,7 @@ mod tests { enable_step_mode: true, mock_spec_call: Some("ci_mock_spec()".into()), enable_mode: false, + available_profiles: vec![], }; let entrypoints = vec![]; @@ -1635,6 +1741,7 @@ mod tests { enable_step_mode: false, mock_spec_call: Some("mock()".into()), enable_mode: false, + available_profiles: vec![], }; let entrypoints = vec![]; @@ -1657,6 +1764,7 @@ mod tests { enable_step_mode: false, mock_spec_call: Some("mock()".into()), enable_mode: false, + available_profiles: vec![], }; let entrypoints = vec![]; @@ -1691,6 +1799,7 @@ mod tests { enable_step_mode: true, mock_spec_call: Some("mock()".into()), enable_mode: false, + available_profiles: vec![], }; let entrypoints = vec![]; @@ -1716,6 +1825,7 @@ mod tests { enable_step_mode: false, mock_spec_call: Some("mock()".into()), enable_mode: true, + available_profiles: vec![], }; let entrypoints = vec![]; @@ -1755,6 +1865,7 @@ mod tests { enable_step_mode: false, mock_spec_call: Some("mock()".into()), enable_mode: false, + available_profiles: vec![], }; let entrypoints = vec![]; @@ -1785,6 +1896,7 @@ mod tests { enable_step_mode: false, mock_spec_call: None, enable_mode: false, + available_profiles: vec![], }; let subcommands = vec![ @@ -1866,4 +1978,88 @@ mod tests { "create subcommand should have owner param in schema" ); } + + #[test] + fn test_generate_cli_with_profile_flag() { + let tool = ToolMeta { + crate_name: "gunbc-sdlc".into(), + tool_name: "sdlc".into(), + description: "SDLC pipeline".into(), + graph_builder_call: "build_sdlc_graph".into(), + graph_builder_args: "".into(), + returns_result: false, + success_port: None, + enable_step_mode: false, + mock_spec_call: Some("mock()".into()), + enable_mode: false, + available_profiles: vec![ + "cloud_run".to_string(), + "local".to_string(), + "unit_test".to_string(), + ], + }; + let entrypoints = vec![]; + + let code = generate_cli(&tool, &entrypoints); + + // Should have profile schema param + assert!( + code.contains("\"profile\""), + "schema should have profile param" + ); + + // Should have profile validation + assert!( + code.contains("valid_profiles"), + "should have profile validation" + ); + assert!( + code.contains("\"cloud_run\""), + "should list cloud_run profile" + ); + assert!( + code.contains("\"local\""), + "should list local profile" + ); + assert!( + code.contains("\"unit_test\""), + "should list unit_test profile" + ); + + // Help should mention profile + assert!( + code.contains("--profile NAME"), + "help should mention --profile flag" + ); + } + + #[test] + fn test_generate_cli_without_profile_flag() { + let tool = ToolMeta { + crate_name: "gunbc-gist".into(), + tool_name: "gist".into(), + description: "Create gist".into(), + graph_builder_call: "build_gist_graph".into(), + graph_builder_args: "".into(), + returns_result: false, + success_port: None, + enable_step_mode: false, + mock_spec_call: Some("mock()".into()), + enable_mode: false, + available_profiles: vec![], + }; + let entrypoints = vec![]; + + let code = generate_cli(&tool, &entrypoints); + + // Should NOT have profile handling + assert!( + !code.contains("valid_profiles"), + "should not have profile validation when no profiles" + ); + assert!( + !code.contains("--profile NAME"), + "help should not mention --profile when no profiles" + ); + } } diff --git a/core/codegen/src/registry.rs b/core/codegen/src/registry.rs index aa0fd073ea1..2df471c0b54 100644 --- a/core/codegen/src/registry.rs +++ b/core/codegen/src/registry.rs @@ -224,6 +224,7 @@ impl ToolDef { enable_step_mode: false, mock_spec_call: None, enable_mode: false, + available_profiles: vec![], }, entrypoints: vec![], custom_import: None, @@ -293,6 +294,15 @@ impl ToolDef { self } + /// Set available profiles for `--profile` enum flag (C20/RT59). + /// + /// When non-empty, the generated CLI accepts `--profile ` to select + /// which interface bindings are active at runtime. + pub fn available_profiles(mut self, profiles: Vec) -> Self { + self.meta.available_profiles = profiles; + self + } + /// Add a subcommand for multi-func dispatch (RT63). pub fn subcommand(mut self, subcmd: SubcommandDef) -> Self { self.subcommands.push(subcmd); diff --git a/core/daglang/daglang-driver/src/lib.rs b/core/daglang/daglang-driver/src/lib.rs index e5d178bf0c5..b2c0081869b 100644 --- a/core/daglang/daglang-driver/src/lib.rs +++ b/core/daglang/daglang-driver/src/lib.rs @@ -63,6 +63,12 @@ pub struct CompileOutput { /// Keys are both qualified (`module.name`) and unqualified (`name`). /// Values are the constant expressions from `data` items. pub data_values: HashMap, + /// Available profile names extracted from `profile` declarations (C20/RT59). + /// + /// When non-empty, the CLI generator can produce a `--profile` enum flag + /// that validates against these names. Profile selection determines which + /// interface bindings are active at runtime. + pub available_profiles: Vec, } impl CompileOutput { @@ -413,6 +419,7 @@ pub fn compile_from_module_graph_with_options( let inferred_entrypoints = daglang_lower::infer_entrypoints(&lowered); let dsl_type_registry = extract_dsl_type_registry(&typed); let data_values = daglang_lower::build_data_values(&typed); + let available_profiles = collect_available_profiles(&typed); let receipt = compute_receipt(&lowered, &emitted, &emit_manifest_path, &source_paths); @@ -427,6 +434,7 @@ pub fn compile_from_module_graph_with_options( dsl_type_registry, receipt, data_values, + available_profiles, }) } @@ -463,6 +471,24 @@ fn collect_pipeline_params(typed: &TypedProject) -> Vec { params } +/// Extract available profile names from the typed project (C20/RT59). +/// +/// Walks all modules looking for `profile` declarations and collects +/// their names. These names become valid values for the `--profile` CLI flag. +fn collect_available_profiles(typed: &TypedProject) -> Vec { + let mut profiles = Vec::new(); + for module in &typed.modules { + for item in &module.ast.items { + if let Item::ProfileDef(def) = &item.node { + profiles.push(def.name.clone()); + } + } + } + profiles.sort(); + profiles.dedup(); + profiles +} + /// Extract a `TypeRegistry` from DSL-defined sum and product types. /// /// Walks all modules in the `TypedProject` and registers: diff --git a/docs/design/transport-primitives.md b/docs/design/transport-primitives.md new file mode 100644 index 00000000000..a8c38f173db --- /dev/null +++ b/docs/design/transport-primitives.md @@ -0,0 +1,364 @@ +# Transport Primitives as DSL Patterns + +**Status**: Proposed +**Lane**: 7 (Transport Domain Modeling) +**Author**: Claude +**Date**: 2026-03-01 + +## Executive Summary + +Transport behaviors (rate limiting, retry, credentials, error parsing) are currently split between: +- **Domain data in Rust** (anti-pattern) — e.g., GitHub's 5000/hour limit hardcoded +- **OS mechanisms in Rust** (correct) — e.g., `TokenBucket`, mutexes, sockets + +This design moves **domain data** into `.dag` while keeping **OS mechanisms** in per-language Target SDKs. + +--- + +## The Core Distinction: "What" vs "How" + +### Domain Policy (What) → belongs in `.dag` + +Questions like: +- "What are GitHub's rate limits?" (5000/hour core, 30/minute search) +- "What status codes trigger retry?" (429, 500, 502, 503, 504) +- "What shape does GCP return errors in?" (`{ "error": { "code": N, "message": "..." } }`) +- "How do we authenticate to this provider?" (Bearer token in `Authorization` header) + +These are **external facts about services** — they belong in `.dag` files alongside service definitions. + +### OS Mechanisms (How) → belongs in Target SDK + +Questions like: +- "How do I pause a thread for 100ms?" (`std::thread::sleep`) +- "How do I atomically decrement a counter?" (`AtomicU64`) +- "How do I share state across threads?" (`Arc>`) +- "How do I open a TCP socket?" (`TcpStream::connect`) + +These are **operating system primitives** — they belong in a Target SDK per language. + +### Why This Matters + +If we tried to express Token Bucket math in `.dag`: +``` +// DON'T DO THIS — turns .dag into a systems language +fn acquire_token() { + let now = Instant::now() // Clock access + let elapsed = now - self.last_refill // Duration math + let tokens = elapsed.as_secs_f64() * rate // Float arithmetic + self.available.fetch_add(tokens, Ordering::Release) // Atomic CAS + ... +} +``` + +We'd have to add mutexes, clocks, atomics, and duration math to `.dag`. That accidentally turns it into C++, defeating its purpose as a high-level orchestration language. + +--- + +## The Architecture: Protobuf/gRPC Pattern + +This follows the same pattern as Protocol Buffers + gRPC: + +``` +.proto (schema) → Generated code → Language runtime +──────────────── ─────────────── ──────────────── +message User { ... } user.pb.go grpc-go library +service UserService { } user_grpc.pb.go (handles sockets, HTTP/2) +``` + +For gunbc: + +``` +.dag (domain policy) → Generated config → Target SDK +──────────────────── ──────────────── ────────────── +rate_limit core { RateLimitMiddleware gunbc-rust-transport + budget: 5000 per hour ::new(5000, 3600) (TokenBucket, Mutex) +} +``` + +**The compiler generates configuration code** that links to a Target SDK. +It does NOT generate the Token Bucket algorithm line-by-line. + +If we want Python tomorrow, we: +1. Write `gunbc-python-transport` once (using `asyncio`) +2. Compiler emits: `RateLimitMiddleware(budget=5000, window_secs=3600)` +3. The Python SDK handles all the async/await mechanics + +--- + +## What Moves to `.dag` + +### 1. Rate Limit Budgets + +**Current (anti-pattern):** Hardcoded in Rust +```rust +// lib/transport/src/rate_limit.rs +const GITHUB_CORE_LIMIT: u64 = 5000; +const GITHUB_SEARCH_LIMIT: u64 = 30; +``` + +**Target:** Declared in service definition +``` +service github.Gist { + config { + rate_limit core { budget: 5000 per hour, burst: 100 } + rate_limit search { budget: 30 per minute } + } +} +``` + +### 2. Retry Policies + +**Current:** Hardcoded retry logic +```rust +const RETRYABLE_STATUSES: &[u16] = &[429, 500, 502, 503, 504]; +``` + +**Target:** Declared per-service or per-operation +``` +service github.Gist { + config { + retry { + max_attempts: 3 + backoff: exponential { base: 100ms, max: 2s, jitter: true } + on: [429, 500, 502, 503, 504] + } + } +} + +operation Create { + retry { max_attempts: 1 } // override: no retry for non-idempotent +} +``` + +### 3. Provider Error Shapes + +**Current (major anti-pattern):** Rust code knows GitHub's error format +```rust +// lib/transport/src/classify.rs +if host.contains("github.com") { + // Parse { "message": ..., "documentation_url": ... } +} +``` + +**Target:** Error shape in service definition +``` +service github.Gist { + config { + error_shape { + message: ".message" + code: ".status" + docs: ".documentation_url" + } + } +} +``` + +### 4. Credential Injection + +**Current:** Auth scheme hardcoded per provider +```rust +fn inject_auth(provider: &str, token: &str) -> Header { + match provider { + "github" => ("Authorization", format!("Bearer {}", token)), + "gcp" => ("Authorization", format!("Bearer {}", token)), + ... + } +} +``` + +**Target:** Credential block in service config +``` +service github.Gist { + config { + credential token { + provider: BearerToken + inject: header("Authorization", "Bearer {token}") + } + } +} +``` + +--- + +## What Stays in Target SDK + +The following remain in `lib/transport/` (Rust) or equivalent per-language SDK: + +| Component | Why it stays | +|-----------|--------------| +| `TokenBucket` struct | Atomic counters, clock access, thread sleep | +| `RetryExecutor` | Loop mechanics, backoff calculation, jitter RNG | +| `CircuitBreaker` | State machine, mutex-protected transitions | +| `CredentialCache` | TTL tracking, thread-safe refresh | +| `HttpClient` | TCP sockets, TLS, HTTP/2 framing | + +The Target SDK is **domain-agnostic**. It doesn't know "GitHub" or "GCP" — it only knows: +- "I was configured with budget=5000, window=3600" +- "I should retry on status codes [429, 500, 502, 503, 504]" +- "I should inject this header with this format" + +--- + +## DSL Syntax + +### `rate_limit` Block + +``` +rate_limit { + budget: per + burst: ? // default: 10% of budget + honor_retry_after: ? // default: true +} +``` + +### `retry` Block + +``` +retry { + max_attempts: + backoff: fixed() + | linear { base: , step: } + | exponential { base: , max: , jitter: } + on: [] + on_network_error: ? // default: true +} +``` + +### `error_shape` Block + +``` +error_shape { + message: + code: ? + details: ? +} +``` + +### `credential` Block + +``` +credential { + provider: BearerToken | ApiKey | GcpWif | AwsSigV4 + inject: header(, ) + | query() + cache_key: ? + refresh_at: ? // default: 80% +} +``` + +--- + +## Compilation Flow + +### Phase 1: Parse & Typecheck + +``` +rate_limit core { budget: 5000 per hour } + ↓ parse +RateLimitDecl { name: "core", budget: 5000, window: Hour, ... } + ↓ typecheck +✓ budget > 0, window is valid unit, name unique within service +``` + +### Phase 2: Lower to IR + +``` +RateLimitDecl + ↓ lower +TransportMiddlewareConfig::RateLimit { + id: "github.Gist.core", + tokens_per_second: 5000.0 / 3600.0, + burst: 100, +} +``` + +### Phase 3: Emit (per target) + +**Rust emit:** +```rust +// Generated: links to gunbc-rust-transport SDK +let rate_limiter = RateLimitMiddleware::new(RateLimitConfig { + id: "github.Gist.core", + tokens_per_second: 1.389, + burst: 100, +}); +``` + +**Go emit (future):** +```go +// Generated: links to gunbc-go-transport SDK +rateLimiter := transport.NewRateLimiter(transport.RateLimitConfig{ + ID: "github.Gist.core", + TokensPerSecond: 1.389, + Burst: 100, +}) +``` + +--- + +## Migration Path + +### Phase 1: DSL Syntax + IR (no runtime changes) + +1. Add grammar rules for `rate_limit`, `retry`, `error_shape`, `credential` +2. Parse into AST nodes +3. Typecheck scopes and references +4. Lower to existing `TransportMiddlewareConfig` IR +5. **Rust runtime continues to interpret** — no behavioral changes + +**Deliverable:** `.dag` files can express transport policy; compiler validates and lowers. + +### Phase 2: Domain Data Migration + +1. Move hardcoded values from Rust to `.dag` service definitions +2. Delete `GITHUB_CORE_LIMIT` constants +3. Delete provider-specific `if host.contains("github.com")` branches +4. Target SDK reads all config from compiled IR + +**Deliverable:** Rust transport code is provider-agnostic. + +### Phase 3: Multi-Target Emit + +1. Emit transport configuration per target language +2. Define `gunbc-go-transport` SDK interface +3. Compiler emits Go code that configures the Go SDK + +**Deliverable:** Same `.dag` file produces working Rust and Go code. + +### Phase 4: Substrate Cleanup + +1. `lib/transport/` becomes a pure Target SDK (no domain knowledge) +2. Delete `classify.rs` provider-specific branches +3. All domain facts live in `.dag` + +**Deliverable:** Complete separation of concerns. + +--- + +## Design Decisions (Pre-Resolved) + +| Decision | Choice | Rationale | +|----------|--------|-----------| +| Rate limit state scope | Per-process | Distributed state requires external store; defer to later | +| Retry jitter | Full jitter (0 to backoff) | AWS recommendation; prevents thundering herd | +| Circuit breaker sharing | Per-service, not per-operation | Most APIs share rate limits across endpoints | +| Credential refresh | Proactive at 80% TTL | Avoids request-time refresh latency | +| Error shape extraction | JSON path syntax | Familiar, handles nested structures | + +## Open Questions (To Resolve Before Implementation) + +1. **Shared state across replicas?** For distributed deployments, rate limiters need coordination. Options: Redis, distributed token bucket, or "best effort" per-replica. *Decision needed before Phase 3.* + +2. **Observability hooks?** Should `.dag` have `on_request`, `on_response`, `on_retry` blocks for metrics/logging? Or is this Target SDK concern? *Decision needed before Phase 2.* + +3. **Test fixtures?** How do `rate_limit` and `retry` interact with `mock` blocks? Should tests auto-skip rate limiting? *Decision needed before Phase 1.* + +--- + +## References + +- `docs/handbook.md` § Compositional Modeling Philosophy +- `docs/design/modeling/annotation-to-dag-modeling.md` — annotation migration tracking +- `SPEC.md` § Transport Model — IR representation +- External: [AWS Exponential Backoff](https://docs.aws.amazon.com/general/latest/gr/api-retries.html) diff --git a/lib/transport/src/classify.rs b/lib/transport/src/classify.rs index 020463c9171..05368ad91ce 100644 --- a/lib/transport/src/classify.rs +++ b/lib/transport/src/classify.rs @@ -8,7 +8,7 @@ use serde_json::Value as JsonValue; use std::collections::HashMap; /// Normalized transport error kind. -#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] pub enum ClassifiedErrorKind { Auth, RateLimit, @@ -134,36 +134,98 @@ fn classify_status( ClassifiedErrorKind::Unknown } -fn provider_error_message(provider: ResponseProvider, body: &JsonValue) -> Option { +/// Provider-specific error details parsed from response body. +#[derive(Debug, Clone, Default)] +pub struct ProviderDiagnostics { + /// Primary error message. + pub message: Option, + /// Error type/code (provider-specific). + pub error_type: Option, + /// HTTP status string (GCP). + pub status: Option, + /// Documentation URL (GitHub). + pub documentation_url: Option, + /// Numeric error code (GCP, OpenAI). + pub code: Option, +} + +impl ProviderDiagnostics { + /// Combine fields into a human-readable message. + pub fn to_message(&self) -> Option { + if let Some(msg) = &self.message { + let mut result = msg.clone(); + if let Some(doc) = &self.documentation_url { + result.push_str(" (see: "); + result.push_str(doc); + result.push(')'); + } + Some(result) + } else { + self.error_type.clone() + } + } +} + +/// Parse provider-specific error shape from response body. +pub fn parse_provider_error(provider: ResponseProvider, body: &JsonValue) -> ProviderDiagnostics { match provider { - ResponseProvider::GitHub => body - .get("message") - .and_then(JsonValue::as_str) - .map(|s| s.to_string()), - ResponseProvider::Gcp => body - .get("error") - .and_then(|v| v.get("message").or_else(|| v.get("status"))) - .and_then(JsonValue::as_str) - .map(|s| s.to_string()), - ResponseProvider::Anthropic => body - .get("error") - .and_then(|v| v.get("message")) - .and_then(JsonValue::as_str) - .map(|s| s.to_string()) - .or_else(|| { - body.get("type") - .and_then(JsonValue::as_str) - .map(|s| s.to_string()) - }), - ResponseProvider::OpenAi => body - .get("error") - .and_then(|v| v.get("message")) + ResponseProvider::GitHub => parse_github_error(body), + ResponseProvider::Gcp => parse_gcp_error(body), + ResponseProvider::Anthropic => parse_anthropic_error(body), + ResponseProvider::OpenAi => parse_openai_error(body), + ResponseProvider::Generic => ProviderDiagnostics::default(), + } +} + +/// GitHub error shape: `{ message, documentation_url }` +fn parse_github_error(body: &JsonValue) -> ProviderDiagnostics { + ProviderDiagnostics { + message: body.get("message").and_then(JsonValue::as_str).map(String::from), + documentation_url: body.get("documentation_url").and_then(JsonValue::as_str).map(String::from), + ..Default::default() + } +} + +/// GCP error shape: `{ error: { code, message, status } }` +fn parse_gcp_error(body: &JsonValue) -> ProviderDiagnostics { + let error = body.get("error"); + ProviderDiagnostics { + message: error.and_then(|e| e.get("message")).and_then(JsonValue::as_str).map(String::from), + status: error.and_then(|e| e.get("status")).and_then(JsonValue::as_str).map(String::from), + code: error.and_then(|e| e.get("code")).and_then(JsonValue::as_i64), + ..Default::default() + } +} + +/// Anthropic error shape: `{ type, error: { type, message } }` +fn parse_anthropic_error(body: &JsonValue) -> ProviderDiagnostics { + let error = body.get("error"); + ProviderDiagnostics { + message: error.and_then(|e| e.get("message")).and_then(JsonValue::as_str).map(String::from), + error_type: error + .and_then(|e| e.get("type")) .and_then(JsonValue::as_str) - .map(|s| s.to_string()), - ResponseProvider::Generic => None, + .map(String::from) + .or_else(|| body.get("type").and_then(JsonValue::as_str).map(String::from)), + ..Default::default() } } +/// OpenAI error shape: `{ error: { message, type, code } }` +fn parse_openai_error(body: &JsonValue) -> ProviderDiagnostics { + let error = body.get("error"); + ProviderDiagnostics { + message: error.and_then(|e| e.get("message")).and_then(JsonValue::as_str).map(String::from), + error_type: error.and_then(|e| e.get("type")).and_then(JsonValue::as_str).map(String::from), + code: error.and_then(|e| e.get("code")).and_then(JsonValue::as_i64), + ..Default::default() + } +} + +fn provider_error_message(provider: ResponseProvider, body: &JsonValue) -> Option { + parse_provider_error(provider, body).to_message() +} + fn message_has_auth_indicator(message: Option<&str>) -> bool { message .map(|m| { @@ -198,6 +260,86 @@ fn parse_retry_after_ms(headers: &HashMap) -> Option { raw.parse::().ok().map(|seconds| seconds * 1000) } +// --------------------------------------------------------------------------- +// Middleware integration +// --------------------------------------------------------------------------- + +use gunbc_ir::transport::TransportResponse; + +/// Default classification policy when none is configured. +fn default_policy() -> ResponseClassification { + ResponseClassification { + provider: ResponseProvider::Generic, + prioritize_auth_errors: true, + parse_provider_error_shapes: true, + } +} + +/// Classify a transport response for middleware decisions. +/// +/// This is the primary entry point for middleware (retry, metrics) to classify +/// responses. It handles the `TransportResponse` enum and dispatches to the +/// appropriate type-specific classifier. +/// +/// Returns `None` for: +/// - Successful HTTP responses (2xx) +/// - Non-HTTP transports (File, Shell, Tcp, Local) +/// +/// # Arguments +/// +/// * `response` - The transport response to classify +/// * `policy` - Optional classification policy; uses a sensible default if None +/// +/// # Example +/// +/// ```ignore +/// use gunbc_lib_transport::classify::{classify_for_middleware, ClassifiedErrorKind}; +/// +/// let classified = classify_for_middleware(&response, None); +/// if let Some(c) = classified { +/// if c.retryable() { +/// // Handle retry +/// } +/// } +/// ``` +pub fn classify_for_middleware( + response: &TransportResponse, + policy: Option<&ResponseClassification>, +) -> Option { + let default = default_policy(); + let policy = policy.unwrap_or(&default); + match response { + TransportResponse::Rest(r) => classify_rest_response(r, policy), + TransportResponse::Http(r) => classify_http_response(r, policy), + // Non-HTTP transports don't have HTTP-style classification + TransportResponse::File(_) + | TransportResponse::Shell(_) + | TransportResponse::Tcp(_) + | TransportResponse::Local(_) => None, + } +} + +/// Extract HTTP status code from a transport response if applicable. +pub fn extract_status_code(response: &TransportResponse) -> Option { + match response { + TransportResponse::Rest(r) => Some(r.status), + TransportResponse::Http(r) => Some(r.status), + _ => None, + } +} + +/// Check if a response indicates success (2xx for HTTP, success field for others). +pub fn is_success(response: &TransportResponse) -> bool { + match response { + TransportResponse::Rest(r) => r.is_success(), + TransportResponse::Http(r) => r.is_success(), + TransportResponse::File(r) => r.success, + TransportResponse::Shell(r) => r.success(), + TransportResponse::Tcp(_) => true, // TCP connections that succeed don't error + TransportResponse::Local(_) => true, + } +} + #[cfg(test)] mod tests { use super::*; @@ -310,4 +452,128 @@ mod tests { assert_eq!(classified.kind, ClassifiedErrorKind::Network); assert!(classified.retryable()); } + + // Middleware integration tests + + #[test] + fn classify_for_middleware_handles_rest_response() { + use gunbc_ir::transport::TransportResponse; + + let rest_response = RestResponse::new(429, serde_json::json!({"message": "rate limit"})); + let response = TransportResponse::Rest(rest_response); + + let classified = classify_for_middleware(&response, None); + assert!(classified.is_some()); + let c = classified.unwrap(); + assert_eq!(c.kind, ClassifiedErrorKind::RateLimit); + assert!(c.retryable()); + } + + #[test] + fn classify_for_middleware_returns_none_for_success() { + use gunbc_ir::transport::TransportResponse; + + let rest_response = RestResponse::new(200, serde_json::json!({"ok": true})); + let response = TransportResponse::Rest(rest_response); + + let classified = classify_for_middleware(&response, None); + assert!(classified.is_none()); + } + + #[test] + fn classify_for_middleware_returns_none_for_non_http() { + use gunbc_ir::transport::{LocalResponse, ShellResponse, TransportResponse}; + + let shell = TransportResponse::Shell(ShellResponse { + stdout: "ok".to_string(), + stderr: "".to_string(), + exit_code: 0, + }); + assert!(classify_for_middleware(&shell, None).is_none()); + + let local = TransportResponse::Local(LocalResponse { + outputs: serde_json::json!({}), + }); + assert!(classify_for_middleware(&local, None).is_none()); + } + + #[test] + fn extract_status_code_works_for_http_types() { + use gunbc_ir::transport::{LocalResponse, TransportResponse}; + + let rest = TransportResponse::Rest(RestResponse::new(404, serde_json::json!({}))); + assert_eq!(extract_status_code(&rest), Some(404)); + + let http = TransportResponse::Http(HttpResponse { + status: 500, + headers: HashMap::new(), + body: "".to_string(), + }); + assert_eq!(extract_status_code(&http), Some(500)); + + let local = TransportResponse::Local(LocalResponse { + outputs: serde_json::json!({}), + }); + assert_eq!(extract_status_code(&local), None); + } + + #[test] + fn is_success_checks_all_transport_types() { + use gunbc_ir::transport::{FileResponse, LocalResponse, ShellResponse, TransportResponse}; + + // REST success + let rest = TransportResponse::Rest(RestResponse::new(200, serde_json::json!({}))); + assert!(is_success(&rest)); + + // REST failure + let rest_fail = TransportResponse::Rest(RestResponse::new(500, serde_json::json!({}))); + assert!(!is_success(&rest_fail)); + + // Shell success + let shell = TransportResponse::Shell(ShellResponse { + stdout: "".to_string(), + stderr: "".to_string(), + exit_code: 0, + }); + assert!(is_success(&shell)); + + // Shell failure + let shell_fail = TransportResponse::Shell(ShellResponse { + stdout: "".to_string(), + stderr: "error".to_string(), + exit_code: 1, + }); + assert!(!is_success(&shell_fail)); + + // File success is determined by the success field + use gunbc_ir::transport::FileOp; + let file_ok = TransportResponse::File(FileResponse { + path: "/tmp/test".to_string(), + operation: FileOp::Read, + success: true, + content: None, + bytes: None, + exists: Some(true), + error: None, + }); + assert!(is_success(&file_ok)); + + // File failure (success=false) should return false + let file_fail = TransportResponse::File(FileResponse { + path: "/tmp/missing".to_string(), + operation: FileOp::Read, + success: false, + content: None, + bytes: None, + exists: Some(false), + error: Some("file not found".to_string()), + }); + assert!(!is_success(&file_fail)); + + // Local always success + let local = TransportResponse::Local(LocalResponse { + outputs: serde_json::json!({}), + }); + assert!(is_success(&local)); + } } diff --git a/lib/transport/src/credential.rs b/lib/transport/src/credential.rs new file mode 100644 index 00000000000..9dbf1c3f9da --- /dev/null +++ b/lib/transport/src/credential.rs @@ -0,0 +1,424 @@ +//! Credential middleware with TTL-aware caching and proactive refresh. +//! +//! Provides credential management for the transport pipeline: +//! - Caches credentials by key to avoid repeated acquisition +//! - Tracks TTL and proactively refreshes at configurable threshold (default 80%) +//! - Thread-safe credential store for concurrent access +//! +//! # Configuration +//! +//! ```ignore +//! CredentialConfig { +//! provider: CredentialProvider::OAuthBearer, +//! injection: CredentialInjection::AuthorizationBearer, +//! cache_key: Some("github-token".to_string()), +//! cache_ttl_ms: None, // Use credential's natural TTL +//! refresh_threshold_pct: 80, // Proactive refresh at 80% of TTL +//! } +//! ``` + +use crate::middleware::{MiddlewareContext, MiddlewareOutcome, TransportMiddleware}; +use gunbc_exec::ExecError; +use gunbc_ir::transport::{CredentialConfig, TransportRequest}; +use gunbc_ir::{AuthScheme, Credential, Secret}; +use std::collections::HashMap; +use std::sync::Mutex; +use std::time::{Duration, Instant, SystemTime}; + +/// Cached credential entry with timing metadata. +#[derive(Debug)] +struct CachedCredential { + credential: Credential, + fetched_at: Instant, + /// Absolute expiry time (from credential's TTL or config override). + expires_at: Option, + /// Total TTL duration computed at fetch time. + /// Used for proactive refresh threshold calculation. + total_ttl: Option, +} + +impl CachedCredential { + fn new(credential: Credential, ttl_override: Option) -> Self { + let fetched_at = Instant::now(); + let now = SystemTime::now(); + + let (expires_at, total_ttl) = if let Some(ttl_ms) = ttl_override { + let ttl = Duration::from_millis(ttl_ms); + (Some(now + ttl), Some(ttl)) + } else if let Some(expiry) = credential.secret().expires_at() { + // Compute TTL from expiry - now + let ttl = expiry.duration_since(now).ok(); + (Some(expiry), ttl) + } else { + (None, None) + }; + + Self { + credential, + fetched_at, + expires_at, + total_ttl, + } + } + + /// Whether this credential is still valid (not expired). + fn is_valid(&self) -> bool { + match self.expires_at { + Some(expiry) => SystemTime::now() < expiry, + None => true, + } + } + + /// Whether this credential should be proactively refreshed. + fn should_refresh(&self, threshold_pct: u8) -> bool { + let Some(total_ttl) = self.total_ttl else { + return false; // No expiry = never refresh + }; + + // threshold_duration is when we should start refreshing + // e.g., 80% threshold on 1 hour TTL = start refreshing after 48 minutes + let threshold_duration = total_ttl * threshold_pct as u32 / 100; + let elapsed = self.fetched_at.elapsed(); + + elapsed > threshold_duration + } +} + +/// Thread-safe credential cache. +#[derive(Debug, Default)] +pub struct CredentialCache { + entries: Mutex>, +} + +impl CredentialCache { + pub fn new() -> Self { + Self::default() + } + + /// Get a cached credential if valid. + pub fn get(&self, key: &str) -> Option { + let entries = self.entries.lock().unwrap(); + entries.get(key).filter(|c| c.is_valid()).map(|c| c.credential.clone()) + } + + /// Check if credential should be refreshed. + pub fn should_refresh(&self, key: &str, threshold_pct: u8) -> bool { + let entries = self.entries.lock().unwrap(); + entries + .get(key) + .map(|c| c.should_refresh(threshold_pct)) + .unwrap_or(true) // Not cached = needs refresh + } + + /// Store a credential with optional TTL override. + pub fn put(&self, key: String, credential: Credential, ttl_override: Option) { + let mut entries = self.entries.lock().unwrap(); + entries.insert(key, CachedCredential::new(credential, ttl_override)); + } + + /// Remove a cached credential. + pub fn remove(&self, key: &str) { + let mut entries = self.entries.lock().unwrap(); + entries.remove(key); + } + + /// Clear all cached credentials. + pub fn clear(&self) { + let mut entries = self.entries.lock().unwrap(); + entries.clear(); + } + + /// Number of cached entries. + pub fn len(&self) -> usize { + self.entries.lock().unwrap().len() + } + + /// Whether cache is empty. + pub fn is_empty(&self) -> bool { + self.len() == 0 + } +} + +/// Credential provider function type. +/// +/// Takes a config and returns a credential (or error). +pub type CredentialProviderFn = + Box Result + Send + Sync>; + +/// Credential middleware. +pub struct CredentialMiddleware { + config: CredentialConfig, + cache: CredentialCache, + provider: Option, +} + +impl CredentialMiddleware { + /// Create middleware with configuration. + /// + /// Without a provider function, the middleware will only work with + /// credentials supplied externally (via context or request auth field). + pub fn new(config: CredentialConfig) -> Self { + Self { + config, + cache: CredentialCache::new(), + provider: None, + } + } + + /// Create middleware with a credential provider function. + pub fn with_provider(config: CredentialConfig, provider: CredentialProviderFn) -> Self { + Self { + config, + cache: CredentialCache::new(), + provider: Some(provider), + } + } + + /// Get or acquire a credential. + fn get_credential(&self) -> Result { + let cache_key = self.config.cache_key.clone().unwrap_or_else(|| { + format!("{:?}:{:?}", self.config.provider, self.config.injection) + }); + + // Check cache + let cached = self.cache.get(&cache_key); + let needs_refresh = self + .cache + .should_refresh(&cache_key, self.config.refresh_threshold_pct); + + // If we have a valid cached credential and don't need refresh, use it + if let Some(ref cred) = cached { + if !needs_refresh { + return Ok(cred.clone()); + } + } + + // Try to acquire new credential (either proactive refresh or initial fetch) + if let Some(provider) = &self.provider { + match provider(&self.config) { + Ok(credential) => { + self.cache + .put(cache_key, credential.clone(), self.config.cache_ttl_ms); + return Ok(credential); + } + Err(e) => { + // If refresh failed but we have a valid cached credential, use it + if let Some(cred) = cached { + return Ok(cred); + } + // No cached credential and refresh failed - propagate error + return Err(e); + } + } + } + + // No provider - can only use cached + if let Some(cred) = cached { + Ok(cred) + } else { + Err(ExecError::new( + "No credential provider configured and no cached credential available", + )) + } + } +} + +impl TransportMiddleware for CredentialMiddleware { + fn pre_request( + &self, + mut request: TransportRequest, + _ctx: &mut MiddlewareContext, + ) -> MiddlewareOutcome { + // Only apply to REST requests + if let TransportRequest::Rest(ref mut rest) = request { + // Check if request already has auth + if rest.auth.is_some() { + // Auth already set, apply it + if let Some(cred) = rest.auth.take() { + cred.apply(rest); + } + return MiddlewareOutcome::Continue(request); + } + + // Check if request requires auth + if !rest.requires_auth { + return MiddlewareOutcome::Continue(request); + } + + // Get credential and apply + match self.get_credential() { + Ok(cred) => { + cred.apply(rest); + MiddlewareOutcome::Continue(request) + } + Err(e) => MiddlewareOutcome::Abort(e), + } + } else { + // Non-REST requests pass through + MiddlewareOutcome::Continue(request) + } + } + + fn name(&self) -> &'static str { + "credential" + } +} + +/// Create a static credential provider for testing. +pub fn static_credential_provider(token: &str) -> CredentialProviderFn { + let token = token.to_string(); + Box::new(move |_config| { + Ok(Credential::new( + Secret::static_value(token.clone()), + AuthScheme::Bearer, + )) + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::middleware::SharedMiddlewareState; + use gunbc_ir::transport::{ + CredentialInjection, CredentialProvider, LocalRequest, RestRequest, + TransportMiddlewareConfig, + }; + use std::sync::Arc; + + fn test_config() -> CredentialConfig { + CredentialConfig { + provider: CredentialProvider::OAuthBearer, + injection: CredentialInjection::AuthorizationBearer, + cache_key: Some("test".to_string()), + cache_ttl_ms: None, + refresh_threshold_pct: 80, + } + } + + #[test] + fn cache_stores_and_retrieves() { + let cache = CredentialCache::new(); + let cred = Credential::new(Secret::static_value("token123"), AuthScheme::Bearer); + + cache.put("key1".to_string(), cred.clone(), None); + + let retrieved = cache.get("key1"); + assert!(retrieved.is_some()); + } + + #[test] + fn cache_returns_none_for_missing_key() { + let cache = CredentialCache::new(); + assert!(cache.get("nonexistent").is_none()); + } + + #[test] + fn cache_removes_entry() { + let cache = CredentialCache::new(); + let cred = Credential::new(Secret::static_value("token"), AuthScheme::Bearer); + + cache.put("key".to_string(), cred, None); + assert!(!cache.is_empty()); + + cache.remove("key"); + assert!(cache.is_empty()); + } + + #[test] + fn middleware_applies_credential_to_rest_request() { + let config = test_config(); + let mw = CredentialMiddleware::with_provider(config, static_credential_provider("test-token")); + + let mw_config = Arc::new(TransportMiddlewareConfig::default()); + let shared = Arc::new(SharedMiddlewareState::new()); + let mut ctx = MiddlewareContext::new("test.op", false, true, mw_config, shared); + + let mut request = RestRequest::get("https://api.example.com"); + request.requires_auth = true; + let request = TransportRequest::Rest(request); + + let outcome = mw.pre_request(request, &mut ctx); + + match outcome { + MiddlewareOutcome::Continue(TransportRequest::Rest(r)) => { + assert!(r.headers.contains_key("Authorization")); + assert!(r.headers["Authorization"].starts_with("Bearer ")); + } + _ => panic!("Expected Continue with Rest request"), + } + } + + #[test] + fn middleware_skips_non_rest_requests() { + let config = test_config(); + let mw = CredentialMiddleware::new(config); + + let mw_config = Arc::new(TransportMiddlewareConfig::default()); + let shared = Arc::new(SharedMiddlewareState::new()); + let mut ctx = MiddlewareContext::new("test.op", false, true, mw_config, shared); + + let request = TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}), + }); + + let outcome = mw.pre_request(request, &mut ctx); + assert!(matches!(outcome, MiddlewareOutcome::Continue(_))); + } + + #[test] + fn middleware_skips_requests_not_requiring_auth() { + let config = test_config(); + let mw = CredentialMiddleware::new(config); + + let mw_config = Arc::new(TransportMiddlewareConfig::default()); + let shared = Arc::new(SharedMiddlewareState::new()); + let mut ctx = MiddlewareContext::new("test.op", false, true, mw_config, shared); + + let mut request = RestRequest::get("https://api.example.com"); + request.requires_auth = false; + let request = TransportRequest::Rest(request); + + let outcome = mw.pre_request(request, &mut ctx); + + match outcome { + MiddlewareOutcome::Continue(TransportRequest::Rest(r)) => { + assert!(!r.headers.contains_key("Authorization")); + } + _ => panic!("Expected Continue with Rest request"), + } + } + + #[test] + fn middleware_caches_credentials() { + let config = test_config(); + let call_count = Arc::new(std::sync::atomic::AtomicU32::new(0)); + let call_count_clone = call_count.clone(); + + let provider: CredentialProviderFn = Box::new(move |_| { + call_count_clone.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + Ok(Credential::new( + Secret::static_value("token"), + AuthScheme::Bearer, + )) + }); + + let mw = CredentialMiddleware::with_provider(config, provider); + + // First call should invoke provider + let _ = mw.get_credential(); + assert_eq!(call_count.load(std::sync::atomic::Ordering::SeqCst), 1); + + // Second call should use cache + let _ = mw.get_credential(); + assert_eq!(call_count.load(std::sync::atomic::Ordering::SeqCst), 1); + } + + #[test] + fn middleware_fails_without_provider_and_cache() { + let config = test_config(); + let mw = CredentialMiddleware::new(config); + + let result = mw.get_credential(); + assert!(result.is_err()); + } +} diff --git a/lib/transport/src/lib.rs b/lib/transport/src/lib.rs index 4e1bea73a55..54bc18bca28 100644 --- a/lib/transport/src/lib.rs +++ b/lib/transport/src/lib.rs @@ -35,8 +35,15 @@ pub mod backend; pub mod classify; pub mod cli; +pub mod credential; pub mod executor; pub mod freshness_policy; +pub mod metrics; +pub mod middleware; +pub mod rate_limit; +pub mod retry; +pub mod pipeline; +pub mod transport_types; pub mod ops; pub mod preflight; @@ -52,6 +59,26 @@ pub use ops::TransportOps; pub use resource_io::TransportIo; +// Middleware infrastructure +pub use classify::{ + classify_for_middleware, classify_rest_response, classify_transport_error, + extract_status_code, is_success, ClassifiedErrorKind, ClassifiedResponse, +}; +pub use metrics::{InMemoryMetricsSink, LogMetricsSink, MetricsMiddleware, MetricsSink, NullMetricsSink}; +pub use middleware::{ + MiddlewareContext, MiddlewareOutcome, PostProcessOutcome, SharedMiddlewareState, + TransportMiddleware, +}; +pub use rate_limit::{RateLimitMiddleware, RateLimitState}; +pub use retry::{CircuitBreaker, CircuitState, RetryMiddleware}; +pub use pipeline::{TransportPipeline, TransportPipelineBuilder}; +pub use credential::{CredentialCache, CredentialMiddleware}; + +// Transport foundation types (TL-0) +pub use transport_types::{ + EndpointBehavior, FailureMode, OperationBehavior, TransportCapabilities, TransportClass, +}; + pub mod system_models; #[cfg(test)] diff --git a/lib/transport/src/metrics.rs b/lib/transport/src/metrics.rs new file mode 100644 index 00000000000..729c04e0dd2 --- /dev/null +++ b/lib/transport/src/metrics.rs @@ -0,0 +1,491 @@ +//! Transport metrics hooks. +//! +//! Provides observability for transport operations: request counts, timing, +//! retry tracking, rate limit headroom, and error distribution. +//! +//! # Usage +//! +//! ```ignore +//! use gunbc_lib_transport::metrics::{MetricsSink, LogMetricsSink, MetricsMiddleware}; +//! +//! let sink = Arc::new(LogMetricsSink::new()); +//! let middleware = MetricsMiddleware::new(sink); +//! ``` + +use crate::classify::ClassifiedErrorKind; +use crate::middleware::{ + MiddlewareContext, MiddlewareOutcome, PostProcessOutcome, TransportMiddleware, +}; +use gunbc_ir::transport::{TransportRequest, TransportResponse}; +use std::collections::HashMap; +use std::sync::{Arc, Mutex}; +use std::time::Instant; + +/// Trait for metrics collection sinks. +/// +/// Implementations can log, emit structured events, push to metrics systems, +/// or do nothing (for tests). +pub trait MetricsSink: Send + Sync { + /// Record that a request is starting. + fn record_request(&self, operation_id: &str, transport_kind: &str); + + /// Record that a response was received. + fn record_response(&self, operation_id: &str, status: Option, duration_ms: u64); + + /// Record that a retry is happening. + fn record_retry(&self, operation_id: &str, attempt: u32, reason: &str); + + /// Record rate limit headroom for a scope. + fn record_rate_limit_headroom(&self, scope: &str, headroom: f64); + + /// Record an error classification. + fn record_error(&self, operation_id: &str, error_kind: ClassifiedErrorKind); +} + +/// No-op metrics sink for tests and when metrics are disabled. +#[derive(Debug, Default)] +pub struct NullMetricsSink; + +impl NullMetricsSink { + pub fn new() -> Self { + Self + } +} + +impl MetricsSink for NullMetricsSink { + fn record_request(&self, _operation_id: &str, _transport_kind: &str) {} + fn record_response(&self, _operation_id: &str, _status: Option, _duration_ms: u64) {} + fn record_retry(&self, _operation_id: &str, _attempt: u32, _reason: &str) {} + fn record_rate_limit_headroom(&self, _scope: &str, _headroom: f64) {} + fn record_error(&self, _operation_id: &str, _error_kind: ClassifiedErrorKind) {} +} + +/// Logging metrics sink that writes to stderr. +#[derive(Debug, Default)] +pub struct LogMetricsSink { + /// Whether to include timestamps in log output. + pub include_timestamps: bool, +} + +impl LogMetricsSink { + pub fn new() -> Self { + Self { + include_timestamps: true, + } + } + + pub fn without_timestamps() -> Self { + Self { + include_timestamps: false, + } + } + + fn timestamp(&self) -> String { + if self.include_timestamps { + use std::time::{SystemTime, UNIX_EPOCH}; + let now = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default(); + format!("[{}.{:03}] ", now.as_secs(), now.subsec_millis()) + } else { + String::new() + } + } +} + +impl MetricsSink for LogMetricsSink { + fn record_request(&self, operation_id: &str, transport_kind: &str) { + eprintln!( + "{}METRIC request op={} transport={}", + self.timestamp(), + operation_id, + transport_kind + ); + } + + fn record_response(&self, operation_id: &str, status: Option, duration_ms: u64) { + let status_str = status.map_or("N/A".to_string(), |s| s.to_string()); + eprintln!( + "{}METRIC response op={} status={} duration_ms={}", + self.timestamp(), + operation_id, + status_str, + duration_ms + ); + } + + fn record_retry(&self, operation_id: &str, attempt: u32, reason: &str) { + eprintln!( + "{}METRIC retry op={} attempt={} reason={}", + self.timestamp(), + operation_id, + attempt, + reason + ); + } + + fn record_rate_limit_headroom(&self, scope: &str, headroom: f64) { + eprintln!( + "{}METRIC rate_limit scope={} headroom={:.2}", + self.timestamp(), + scope, + headroom + ); + } + + fn record_error(&self, operation_id: &str, error_kind: ClassifiedErrorKind) { + eprintln!( + "{}METRIC error op={} kind={:?}", + self.timestamp(), + operation_id, + error_kind + ); + } +} + +/// In-memory metrics collector for testing. +#[derive(Debug, Default)] +pub struct InMemoryMetricsSink { + requests: Mutex>, + responses: Mutex, u64)>>, + retries: Mutex>, + rate_limits: Mutex>, + errors: Mutex>, +} + +impl InMemoryMetricsSink { + pub fn new() -> Self { + Self::default() + } + + pub fn request_count(&self) -> usize { + self.requests.lock().unwrap().len() + } + + pub fn response_count(&self) -> usize { + self.responses.lock().unwrap().len() + } + + pub fn retry_count(&self) -> usize { + self.retries.lock().unwrap().len() + } + + pub fn error_count(&self) -> usize { + self.errors.lock().unwrap().len() + } + + pub fn total_duration_ms(&self) -> u64 { + self.responses + .lock() + .unwrap() + .iter() + .map(|(_, _, d)| d) + .sum() + } + + pub fn errors_by_kind(&self) -> HashMap { + let errors = self.errors.lock().unwrap(); + let mut counts = HashMap::new(); + for (_, kind) in errors.iter() { + *counts.entry(*kind).or_insert(0) += 1; + } + counts + } +} + +impl MetricsSink for InMemoryMetricsSink { + fn record_request(&self, operation_id: &str, transport_kind: &str) { + self.requests + .lock() + .unwrap() + .push((operation_id.to_string(), transport_kind.to_string())); + } + + fn record_response(&self, operation_id: &str, status: Option, duration_ms: u64) { + self.responses + .lock() + .unwrap() + .push((operation_id.to_string(), status, duration_ms)); + } + + fn record_retry(&self, operation_id: &str, attempt: u32, reason: &str) { + self.retries + .lock() + .unwrap() + .push((operation_id.to_string(), attempt, reason.to_string())); + } + + fn record_rate_limit_headroom(&self, scope: &str, headroom: f64) { + self.rate_limits + .lock() + .unwrap() + .push((scope.to_string(), headroom)); + } + + fn record_error(&self, operation_id: &str, error_kind: ClassifiedErrorKind) { + self.errors + .lock() + .unwrap() + .push((operation_id.to_string(), error_kind)); + } +} + +/// Per-request timing state. +struct RequestTiming { + start: Instant, +} + +/// Metrics middleware that records request/response timing and counts. +pub struct MetricsMiddleware { + sink: Arc, + /// Active request timings, keyed by unique request_id to avoid collision. + timings: Mutex>, +} + +impl MetricsMiddleware { + pub fn new(sink: Arc) -> Self { + Self { + sink, + timings: Mutex::new(HashMap::new()), + } + } +} + +impl TransportMiddleware for MetricsMiddleware { + fn pre_request( + &self, + request: TransportRequest, + ctx: &mut MiddlewareContext, + ) -> MiddlewareOutcome { + let transport_kind = transport_kind_str(&request); + self.sink.record_request(&ctx.operation_id, transport_kind); + + // Store timing for this request, keyed by unique request_id + let timing = RequestTiming { + start: Instant::now(), + }; + self.timings + .lock() + .unwrap() + .insert(ctx.request_id, timing); + + MiddlewareOutcome::Continue(request) + } + + fn post_response( + &self, + _request: &TransportRequest, + response: TransportResponse, + ctx: &mut MiddlewareContext, + ) -> PostProcessOutcome { + let duration_ms = self + .timings + .lock() + .unwrap() + .remove(&ctx.request_id) + .map(|t| t.start.elapsed().as_millis() as u64) + .unwrap_or(0); + + let status = extract_status(&response); + self.sink + .record_response(&ctx.operation_id, status, duration_ms); + + PostProcessOutcome::Complete(response) + } + + fn on_error( + &self, + _request: &TransportRequest, + error: gunbc_exec::ExecError, + ctx: &mut MiddlewareContext, + ) -> PostProcessOutcome { + // Clean up timing state + self.timings.lock().unwrap().remove(&ctx.request_id); + + // Don't record synthetic cleanup errors as real failures + let error_msg = error.to_string(); + if error_msg.contains("pipeline cleanup") { + return PostProcessOutcome::Abort(error); + } + + // Classify error based on message content + let error_kind = classify_exec_error(&error_msg); + self.sink.record_error(&ctx.operation_id, error_kind); + + PostProcessOutcome::Abort(error) + } + + fn name(&self) -> &'static str { + "metrics" + } +} + +/// Extract transport kind as a string for logging. +fn transport_kind_str(request: &TransportRequest) -> &'static str { + match request { + TransportRequest::Rest(_) => "rest", + TransportRequest::Http(_) => "http", + TransportRequest::File(_) => "file", + TransportRequest::Tcp(_) => "tcp", + TransportRequest::Shell(_) => "shell", + TransportRequest::Local(_) => "local", + } +} + +/// Extract HTTP status from response if applicable. +fn extract_status(response: &TransportResponse) -> Option { + match response { + TransportResponse::Rest(r) => Some(r.status), + TransportResponse::Http(r) => Some(r.status), + _ => None, + } +} + +/// Classify an execution error based on message content. +/// +/// This is a best-effort heuristic for errors that don't have structured +/// classification (e.g., ExecError from transport failures). +fn classify_exec_error(message: &str) -> ClassifiedErrorKind { + let lower = message.to_ascii_lowercase(); + + // Auth errors + if lower.contains("auth") + || lower.contains("credential") + || lower.contains("unauthorized") + || lower.contains("forbidden") + || lower.contains("invalid api key") + || lower.contains("token") + { + return ClassifiedErrorKind::Auth; + } + + // Rate limit errors + if lower.contains("rate limit") || lower.contains("too many requests") || lower.contains("429") + { + return ClassifiedErrorKind::RateLimit; + } + + // Client errors (config, serialization, validation) + if lower.contains("invalid") + || lower.contains("missing") + || lower.contains("config") + || lower.contains("serializ") + || lower.contains("deserializ") + || lower.contains("parse") + { + return ClassifiedErrorKind::Client; + } + + // Server errors + if lower.contains("server") + || lower.contains("internal error") + || lower.contains("5xx") + || lower.contains("500") + { + return ClassifiedErrorKind::Server; + } + + // Default to network for connection/timeout issues + ClassifiedErrorKind::Network +} + +#[cfg(test)] +mod tests { + use super::*; + use gunbc_ir::transport::{LocalRequest, LocalResponse, TransportMiddlewareConfig}; + use std::thread; + use std::time::Duration; + + #[test] + fn null_sink_compiles_and_runs() { + let sink = NullMetricsSink::new(); + sink.record_request("test.op", "rest"); + sink.record_response("test.op", Some(200), 100); + sink.record_retry("test.op", 2, "rate limit"); + sink.record_rate_limit_headroom("github:core", 0.5); + sink.record_error("test.op", ClassifiedErrorKind::RateLimit); + } + + #[test] + fn in_memory_sink_counts_correctly() { + let sink = InMemoryMetricsSink::new(); + sink.record_request("op1", "rest"); + sink.record_request("op2", "shell"); + sink.record_response("op1", Some(200), 50); + sink.record_response("op2", None, 100); + sink.record_retry("op1", 2, "server error"); + sink.record_error("op1", ClassifiedErrorKind::Server); + sink.record_error("op2", ClassifiedErrorKind::Server); + + assert_eq!(sink.request_count(), 2); + assert_eq!(sink.response_count(), 2); + assert_eq!(sink.retry_count(), 1); + assert_eq!(sink.total_duration_ms(), 150); + assert_eq!(sink.error_count(), 2); + assert_eq!( + sink.errors_by_kind().get(&ClassifiedErrorKind::Server), + Some(&2) + ); + } + + #[test] + fn metrics_middleware_records_timing() { + let sink = Arc::new(InMemoryMetricsSink::new()); + let mw = MetricsMiddleware::new(sink.clone()); + let config = Arc::new(TransportMiddlewareConfig::default()); + let shared = Arc::new(crate::middleware::SharedMiddlewareState::new()); + let mut ctx = + crate::middleware::MiddlewareContext::new("test.op", false, true, config, shared); + + let request = TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}), + }); + + // Pre-request + let outcome = mw.pre_request(request, &mut ctx); + assert!(matches!(outcome, MiddlewareOutcome::Continue(_))); + assert_eq!(sink.request_count(), 1); + + // Simulate some work + thread::sleep(Duration::from_millis(10)); + + // Post-response + let response = TransportResponse::Local(LocalResponse { + outputs: serde_json::json!({}), + }); + let outcome = mw.post_response(&TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}), + }), response, &mut ctx); + assert!(matches!(outcome, PostProcessOutcome::Complete(_))); + assert_eq!(sink.response_count(), 1); + assert!(sink.total_duration_ms() >= 10); + } + + #[test] + fn transport_kind_str_matches_variants() { + use gunbc_ir::transport::*; + + assert_eq!( + transport_kind_str(&TransportRequest::Rest(RestRequest::get("https://example.com"))), + "rest" + ); + assert_eq!( + transport_kind_str(&TransportRequest::Shell(ShellRequest { + command: "echo".to_string(), + args: vec![], + cwd: None, + env: HashMap::new(), + stdin: None, + timeout_ms: None, + passthrough: false, + })), + "shell" + ); + assert_eq!( + transport_kind_str(&TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}) + })), + "local" + ); + } +} diff --git a/lib/transport/src/middleware/mod.rs b/lib/transport/src/middleware/mod.rs new file mode 100644 index 00000000000..7fce97c31ed --- /dev/null +++ b/lib/transport/src/middleware/mod.rs @@ -0,0 +1,203 @@ +//! Transport middleware infrastructure. +//! +//! Middleware layers intercept transport requests/responses to provide cross-cutting +//! concerns: rate limiting, retry, credentials, metrics. Each middleware implements +//! `TransportMiddleware` and composes via `TransportPipeline`. +//! +//! # Pipeline Order (standard pipeline) +//! +//! ```text +//! request → metrics → retry → rate_limit → execute → rate_limit → retry → metrics → response +//! ``` +//! +//! Pre-request flows outer→inner; post-response flows inner→outer. +//! The credential layer is available but not included in the standard pipeline. + +use gunbc_exec::ExecError; +use gunbc_ir::transport::{TransportMiddlewareConfig, TransportRequest, TransportResponse}; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::sync::Arc; + +/// Global counter for generating unique request IDs. +static REQUEST_ID_COUNTER: AtomicU64 = AtomicU64::new(1); + +/// Context passed through the middleware chain. +/// +/// Carries request metadata, operation info, and per-request state that middleware +/// layers can read and modify. +#[derive(Debug, Clone)] +pub struct MiddlewareContext { + /// Unique request ID for correlating timing data across concurrent requests. + /// Prevents HashMap collision when multiple requests share the same operation_id. + pub request_id: u64, + /// Operation identifier for logging/metrics (e.g., "github.gist.create"). + pub operation_id: String, + /// Whether the operation is idempotent (safe to retry without side effects). + pub idempotent: bool, + /// Whether the operation is read-only (no server-side state change). + pub readonly: bool, + /// Middleware configuration from IR. + pub config: Arc, + /// Attempt number (1 for first attempt, incremented on retry). + pub attempt: u32, + /// Shared state for cross-request coordination (rate limits, circuit breaker). + pub shared_state: Arc, +} + +impl MiddlewareContext { + /// Create a new context for an operation. + pub fn new( + operation_id: impl Into, + idempotent: bool, + readonly: bool, + config: Arc, + shared_state: Arc, + ) -> Self { + Self { + request_id: REQUEST_ID_COUNTER.fetch_add(1, Ordering::Relaxed), + operation_id: operation_id.into(), + idempotent, + readonly, + config, + attempt: 1, + shared_state, + } + } + + /// Check if the operation is safe to retry (idempotent or readonly). + pub fn retry_safe(&self) -> bool { + self.idempotent || self.readonly + } +} + +/// Shared state across middleware invocations. +/// +/// Holds stateful middleware components like rate limit buckets and circuit breaker +/// state that must be shared across concurrent requests. +#[derive(Debug, Default)] +pub struct SharedMiddlewareState { + // Rate limit state is added by TL-1 + // Circuit breaker state is added by TL-2 +} + +impl SharedMiddlewareState { + /// Create new shared state. + pub fn new() -> Self { + Self::default() + } +} + +/// Result of middleware pre-request processing. +#[derive(Debug)] +pub enum MiddlewareOutcome { + /// Continue to next middleware or execute request. + Continue(TransportRequest), + /// Short-circuit with immediate response (cached, rate-limited waiting, etc.). + ShortCircuit(TransportResponse), + /// Abort with error (circuit open, auth failed, etc.). + Abort(ExecError), +} + +/// Result of middleware post-response processing. +#[derive(Debug)] +pub enum PostProcessOutcome { + /// Return response as-is to caller. + Complete(TransportResponse), + /// Retry the request after delay. + Retry { + /// Delay before retry in milliseconds. + delay_ms: u64, + /// Reason for retry (for logging/metrics). + reason: String, + }, + /// Abort with error. + Abort(ExecError), +} + +/// Transport middleware layer trait. +/// +/// Middleware intercepts requests before execution and responses after execution. +/// Each method has a default pass-through implementation. +pub trait TransportMiddleware: Send + Sync { + /// Process request before execution. + /// + /// Can modify the request, short-circuit with an immediate response, + /// or abort with an error. + fn pre_request( + &self, + request: TransportRequest, + _ctx: &mut MiddlewareContext, + ) -> MiddlewareOutcome { + MiddlewareOutcome::Continue(request) + } + + /// Process response after execution. + /// + /// Can transform the response, request retry, or abort with an error. + fn post_response( + &self, + _request: &TransportRequest, + response: TransportResponse, + _ctx: &mut MiddlewareContext, + ) -> PostProcessOutcome { + PostProcessOutcome::Complete(response) + } + + /// Handle execution error. + /// + /// Can transform the error, request retry (for network errors), or pass through. + fn on_error( + &self, + _request: &TransportRequest, + error: ExecError, + _ctx: &mut MiddlewareContext, + ) -> PostProcessOutcome { + PostProcessOutcome::Abort(error) + } + + /// Middleware name for logging/metrics. + fn name(&self) -> &'static str; +} + +#[cfg(test)] +mod tests { + use super::*; + use gunbc_ir::transport::LocalRequest; + + struct PassthroughMiddleware; + + impl TransportMiddleware for PassthroughMiddleware { + fn name(&self) -> &'static str { + "passthrough" + } + } + + #[test] + fn default_middleware_passes_through() { + let mw = PassthroughMiddleware; + let config = Arc::new(TransportMiddlewareConfig::default()); + let shared = Arc::new(SharedMiddlewareState::new()); + let mut ctx = MiddlewareContext::new("test.op", false, true, config, shared); + + let request = TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}), + }); + let outcome = mw.pre_request(request.clone(), &mut ctx); + assert!(matches!(outcome, MiddlewareOutcome::Continue(_))); + } + + #[test] + fn context_retry_safe_checks_idempotent_or_readonly() { + let config = Arc::new(TransportMiddlewareConfig::default()); + let shared = Arc::new(SharedMiddlewareState::new()); + + let ctx = MiddlewareContext::new("op1", false, false, config.clone(), shared.clone()); + assert!(!ctx.retry_safe()); + + let ctx = MiddlewareContext::new("op2", true, false, config.clone(), shared.clone()); + assert!(ctx.retry_safe()); + + let ctx = MiddlewareContext::new("op3", false, true, config.clone(), shared.clone()); + assert!(ctx.retry_safe()); + } +} diff --git a/lib/transport/src/pipeline.rs b/lib/transport/src/pipeline.rs new file mode 100644 index 00000000000..32f74dfc6db --- /dev/null +++ b/lib/transport/src/pipeline.rs @@ -0,0 +1,549 @@ +//! Transport middleware pipeline composition. +//! +//! Composes middleware layers into a pipeline that wraps transport execution. +//! Middleware is applied in order: outer layers see the request first on the +//! way in, and the response first on the way out. +//! +//! # Pipeline Order (standard pipeline) +//! +//! ```text +//! Request flow: metrics → retry → rate_limit → execute +//! Response flow: execute → rate_limit → retry → metrics +//! ``` +//! +//! Note: The credential layer is NOT included in the standard pipeline. +//! Use `.layer(Arc::new(CredentialMiddleware::new(config)))` to add it. +//! +//! # Usage +//! +//! ```ignore +//! use gunbc_lib_transport::pipeline::TransportPipelineBuilder; +//! +//! let pipeline = TransportPipelineBuilder::standard(&middleware_config) +//! .with_metrics_sink(Arc::new(LogMetricsSink::new())) +//! .build(); +//! +//! let response = pipeline.execute(request, "my.operation", true, false)?; +//! ``` + +use crate::metrics::{MetricsMiddleware, MetricsSink, NullMetricsSink}; +use crate::middleware::{ + MiddlewareContext, MiddlewareOutcome, PostProcessOutcome, SharedMiddlewareState, + TransportMiddleware, +}; +use crate::rate_limit::RateLimitMiddleware; +use crate::retry::RetryMiddleware; +use gunbc_exec::{ExecError, IntoExecResult}; +use gunbc_ir::transport::{TransportMiddlewareConfig, TransportRequest, TransportResponse}; +use std::sync::Arc; + +/// Executor function type for actual transport execution. +pub type ExecutorFn = + Box Result + Send + Sync>; + +/// Builder for constructing middleware pipelines. +pub struct TransportPipelineBuilder { + layers: Vec>, + config: TransportMiddlewareConfig, + metrics_sink: Option>, + executor: Option, +} + +impl TransportPipelineBuilder { + /// Create a new empty pipeline builder. + pub fn new() -> Self { + Self { + layers: Vec::new(), + config: TransportMiddlewareConfig::default(), + metrics_sink: None, + executor: None, + } + } + + /// Set the middleware configuration. + pub fn with_config(mut self, config: TransportMiddlewareConfig) -> Self { + self.config = config; + self + } + + /// Add a custom middleware layer. + /// + /// Layers are applied in order: first added is outermost. + pub fn layer(mut self, middleware: Arc) -> Self { + self.layers.push(middleware); + self + } + + /// Set a custom metrics sink. + pub fn with_metrics_sink(mut self, sink: Arc) -> Self { + self.metrics_sink = Some(sink); + self + } + + /// Set a custom executor function. + /// + /// By default, the pipeline will use the crate's internal executor. + pub fn with_executor(mut self, executor: ExecutorFn) -> Self { + self.executor = Some(executor); + self + } + + /// Build a standard pipeline based on the middleware configuration. + /// + /// Creates middleware layers based on what's configured: + /// - Always adds metrics (outermost) + /// - Adds rate_limit if configured + /// - Adds retry if configured + pub fn standard(config: &TransportMiddlewareConfig) -> Self { + let mut builder = Self::new().with_config(config.clone()); + + // Metrics is always added (outermost) + // Will be added in build() with the configured sink + + // Retry is outer to rate_limit so that rate_limit sees 429 responses first, + // updates its pause_until state from Retry-After headers, then passes + // the response to retry middleware for the actual retry decision. + if let Some(retry_config) = &config.retry { + builder = builder.layer(Arc::new(RetryMiddleware::new(retry_config.clone()))); + } + + // Rate limit is inner - it sees responses before retry and can extract + // Retry-After headers to update its internal state + if let Some(rate_config) = &config.rate_limit { + builder = builder.layer(Arc::new(RateLimitMiddleware::new(rate_config.clone()))); + } + + builder + } + + /// Build the pipeline. + pub fn build(self) -> TransportPipeline { + let mut layers = Vec::new(); + + // Metrics is always outermost + let metrics_sink = self + .metrics_sink + .unwrap_or_else(|| Arc::new(NullMetricsSink::new())); + layers.push(Arc::new(MetricsMiddleware::new(metrics_sink)) as Arc); + + // Then the configured layers + layers.extend(self.layers); + + TransportPipeline { + layers, + config: Arc::new(self.config), + shared_state: Arc::new(SharedMiddlewareState::new()), + executor: self.executor, + } + } +} + +impl Default for TransportPipelineBuilder { + fn default() -> Self { + Self::new() + } +} + +/// Composed middleware pipeline for transport execution. +pub struct TransportPipeline { + layers: Vec>, + config: Arc, + shared_state: Arc, + executor: Option, +} + +impl TransportPipeline { + /// Execute a request through the middleware pipeline. + /// + /// # Arguments + /// + /// * `request` - The transport request to execute + /// * `operation_id` - Identifier for logging/metrics + /// * `idempotent` - Whether the operation is idempotent + /// * `readonly` - Whether the operation is read-only + pub fn execute( + &self, + request: TransportRequest, + operation_id: &str, + idempotent: bool, + readonly: bool, + ) -> Result { + let mut ctx = MiddlewareContext::new( + operation_id, + idempotent, + readonly, + self.config.clone(), + self.shared_state.clone(), + ); + + self.execute_with_retry(request, &mut ctx) + } + + /// Execute with retry loop. + fn execute_with_retry( + &self, + request: TransportRequest, + ctx: &mut MiddlewareContext, + ) -> Result { + let max_attempts = self + .config + .retry + .as_ref() + .map_or(1, |r| r.max_attempts); + + loop { + let result = self.execute_once(request.clone(), ctx); + + match result { + Ok(PostProcessOutcome::Complete(response)) => return Ok(response), + Ok(PostProcessOutcome::Retry { delay_ms, reason }) => { + if ctx.attempt >= max_attempts { + return Err(ExecError::new(format!( + "Max retry attempts ({}) exceeded: {}", + max_attempts, reason + ))); + } + ctx.attempt += 1; + std::thread::sleep(std::time::Duration::from_millis(delay_ms)); + // Continue loop + } + Ok(PostProcessOutcome::Abort(error)) => return Err(error), + Err(error) => return Err(error), + } + } + } + + /// Execute request once through the pipeline. + /// + /// Tracks which layers have run `pre_request` to ensure proper cleanup + /// when inner layers abort or request retry. + fn execute_once( + &self, + request: TransportRequest, + ctx: &mut MiddlewareContext, + ) -> Result { + // Track how many layers successfully ran pre_request + let mut layers_completed_pre = 0; + // Keep a clone for cleanup in case of early exit + let request_for_cleanup = request.clone(); + let mut current_request = request; + + // Pre-request processing (outer to inner) + for layer in &self.layers { + match layer.pre_request(current_request, ctx) { + MiddlewareOutcome::Continue(req) => { + current_request = req; + layers_completed_pre += 1; + } + MiddlewareOutcome::ShortCircuit(response) => { + // Cleanup layers that already ran pre_request + self.cleanup_layers(layers_completed_pre, &request_for_cleanup, ctx); + return Ok(PostProcessOutcome::Complete(response)); + } + MiddlewareOutcome::Abort(error) => { + // Cleanup layers that already ran pre_request + self.cleanup_layers(layers_completed_pre, &request_for_cleanup, ctx); + return Ok(PostProcessOutcome::Abort(error)); + } + } + } + + // Use the final transformed request for execution + let request = current_request; + + // Execute the actual transport + let result = if let Some(executor) = &self.executor { + executor(&request) + } else { + // Use crate-internal executor + crate::backend::execute_transport_with_backend(&request) + .exec_context("transport error") + }; + + // Post-process result + match result { + Ok(mut response) => { + // Post-response processing (inner to outer) + // All layers ran pre_request, so process all in reverse + for (idx, layer) in self.layers.iter().enumerate().rev() { + match layer.post_response(&request, response, ctx) { + PostProcessOutcome::Complete(resp) => response = resp, + outcome @ PostProcessOutcome::Retry { .. } => { + // Clean up remaining outer layers that haven't seen post_response + self.cleanup_layers(idx, &request, ctx); + return Ok(outcome); + } + outcome @ PostProcessOutcome::Abort(_) => { + // Clean up remaining outer layers + self.cleanup_layers(idx, &request, ctx); + return Ok(outcome); + } + } + } + Ok(PostProcessOutcome::Complete(response)) + } + Err(error) => { + // Error processing (inner to outer) + let mut outcome = PostProcessOutcome::Abort(error); + for (idx, layer) in self.layers.iter().enumerate().rev() { + if let PostProcessOutcome::Abort(e) = outcome { + outcome = layer.on_error(&request, e, ctx); + } else { + // Inner layer transformed error to Retry/Complete + // Clean up remaining outer layers (indices 0..=idx) + self.cleanup_layers(idx + 1, &request, ctx); + break; + } + } + Ok(outcome) + } + } + } + + /// Clean up layers that ran pre_request but won't see post_response. + /// + /// Calls on_error with a synthetic cleanup error for the first N layers + /// (outer layers) that successfully ran pre_request. + fn cleanup_layers( + &self, + count: usize, + request: &TransportRequest, + ctx: &mut MiddlewareContext, + ) { + // Create a synthetic cleanup error - layers use this to clean up state + let cleanup_error = ExecError::new("pipeline cleanup (request did not complete)"); + + // Call on_error for layers 0..count in reverse order (inner to outer cleanup) + for layer in self.layers[..count].iter().rev() { + // Ignore the outcome - we're just cleaning up + let _ = layer.on_error(request, cleanup_error.clone(), ctx); + } + } + + /// Get the number of middleware layers. + pub fn layer_count(&self) -> usize { + self.layers.len() + } + + /// Get the middleware layer names. + pub fn layer_names(&self) -> Vec<&'static str> { + self.layers.iter().map(|l| l.name()).collect() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::metrics::InMemoryMetricsSink; + use gunbc_ir::transport::{ + LocalRequest, LocalResponse, RateLimitAlgorithm, RateLimitConfig, RetryBackoff, + RetryConfig, + }; + + fn mock_executor() -> ExecutorFn { + Box::new(|_req| { + Ok(TransportResponse::Local(LocalResponse { + outputs: serde_json::json!({"result": "ok"}), + })) + }) + } + + fn failing_executor(fail_count: Arc) -> ExecutorFn { + Box::new(move |_req| { + let count = fail_count.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + if count < 2 { + Err(ExecError::new("transient failure")) + } else { + Ok(TransportResponse::Local(LocalResponse { + outputs: serde_json::json!({"result": "ok"}), + })) + } + }) + } + + use std::sync::atomic::AtomicU32; + + #[test] + fn empty_pipeline_executes_directly() { + let pipeline = TransportPipelineBuilder::new() + .with_executor(mock_executor()) + .build(); + + let request = TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}), + }); + + let response = pipeline.execute(request, "test.op", false, true); + assert!(response.is_ok()); + } + + #[test] + fn standard_pipeline_adds_metrics() { + let config = TransportMiddlewareConfig::default(); + let pipeline = TransportPipelineBuilder::standard(&config) + .with_executor(mock_executor()) + .build(); + + assert!(pipeline.layer_names().contains(&"metrics")); + } + + #[test] + fn standard_pipeline_adds_rate_limit_when_configured() { + let config = TransportMiddlewareConfig { + rate_limit: Some(RateLimitConfig { + scope_key: "test".to_string(), + algorithm: RateLimitAlgorithm::TokenBucket, + max_burst: 10, + sustained_per_minute: 60, + honor_retry_after: true, + }), + ..Default::default() + }; + + let pipeline = TransportPipelineBuilder::standard(&config) + .with_executor(mock_executor()) + .build(); + + assert!(pipeline.layer_names().contains(&"rate_limit")); + } + + #[test] + fn standard_pipeline_adds_retry_when_configured() { + let config = TransportMiddlewareConfig { + retry: Some(RetryConfig { + max_attempts: 3, + base_delay_ms: 10, + max_delay_ms: 100, + backoff: RetryBackoff::Fixed, + retry_statuses: vec![500], + retry_network_errors: true, + require_idempotent_or_readonly: true, + circuit_breaker: None, + }), + ..Default::default() + }; + + let pipeline = TransportPipelineBuilder::standard(&config) + .with_executor(mock_executor()) + .build(); + + assert!(pipeline.layer_names().contains(&"retry")); + } + + #[test] + fn pipeline_records_metrics() { + let sink = Arc::new(InMemoryMetricsSink::new()); + let pipeline = TransportPipelineBuilder::new() + .with_metrics_sink(sink.clone()) + .with_executor(mock_executor()) + .build(); + + let request = TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}), + }); + + let _ = pipeline.execute(request, "test.op", false, true); + + assert_eq!(sink.request_count(), 1); + assert_eq!(sink.response_count(), 1); + } + + #[test] + fn pipeline_retries_on_error() { + let config = TransportMiddlewareConfig { + retry: Some(RetryConfig { + max_attempts: 5, + base_delay_ms: 1, // Minimal delay for test + max_delay_ms: 10, + backoff: RetryBackoff::Fixed, + retry_statuses: vec![500], + retry_network_errors: true, + require_idempotent_or_readonly: true, + circuit_breaker: None, + }), + ..Default::default() + }; + + let fail_count = Arc::new(AtomicU32::new(0)); + let pipeline = TransportPipelineBuilder::standard(&config) + .with_executor(failing_executor(fail_count.clone())) + .build(); + + let request = TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}), + }); + + // Should succeed after retries (idempotent = true) + let response = pipeline.execute(request, "test.op", true, false); + assert!(response.is_ok()); + + // Should have tried 3 times (2 failures + 1 success) + assert_eq!(fail_count.load(std::sync::atomic::Ordering::SeqCst), 3); + } + + #[test] + fn pipeline_respects_idempotency_requirement() { + let config = TransportMiddlewareConfig { + retry: Some(RetryConfig { + max_attempts: 5, + base_delay_ms: 1, + max_delay_ms: 10, + backoff: RetryBackoff::Fixed, + retry_statuses: vec![500], + retry_network_errors: true, + require_idempotent_or_readonly: true, + circuit_breaker: None, + }), + ..Default::default() + }; + + let fail_count = Arc::new(AtomicU32::new(0)); + let pipeline = TransportPipelineBuilder::standard(&config) + .with_executor(failing_executor(fail_count.clone())) + .build(); + + let request = TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}), + }); + + // Non-idempotent should NOT retry + let response = pipeline.execute(request, "test.op", false, false); + assert!(response.is_err()); + + // Should have only tried once + assert_eq!(fail_count.load(std::sync::atomic::Ordering::SeqCst), 1); + } + + #[test] + fn layer_names_returns_correct_order() { + let config = TransportMiddlewareConfig { + rate_limit: Some(RateLimitConfig { + scope_key: "test".to_string(), + algorithm: RateLimitAlgorithm::TokenBucket, + max_burst: 10, + sustained_per_minute: 60, + honor_retry_after: true, + }), + retry: Some(RetryConfig { + max_attempts: 3, + base_delay_ms: 10, + max_delay_ms: 100, + backoff: RetryBackoff::Fixed, + retry_statuses: vec![500], + retry_network_errors: true, + require_idempotent_or_readonly: true, + circuit_breaker: None, + }), + ..Default::default() + }; + + let pipeline = TransportPipelineBuilder::standard(&config) + .with_executor(mock_executor()) + .build(); + + let names = pipeline.layer_names(); + // Order: metrics (outermost) → retry → rate_limit (innermost) + // Rate limit is inner to see 429 responses before retry consumes them + assert_eq!(names, vec!["metrics", "retry", "rate_limit"]); + } +} diff --git a/lib/transport/src/rate_limit.rs b/lib/transport/src/rate_limit.rs new file mode 100644 index 00000000000..51e3e9accf9 --- /dev/null +++ b/lib/transport/src/rate_limit.rs @@ -0,0 +1,530 @@ +//! Rate limit middleware. +//! +//! Provides rate limiting using token bucket or sliding window algorithms. +//! Rate limits are tracked per scope (e.g., "github:core", "github:search") +//! to respect provider-specific quotas. +//! +//! # Configuration +//! +//! Rate limiting is configured via `RateLimitConfig` from the IR: +//! +//! ```ignore +//! RateLimitConfig { +//! scope_key: "github:core", +//! algorithm: RateLimitAlgorithm::TokenBucket, +//! max_burst: 20, +//! sustained_per_minute: 83, // ~5000/hour +//! honor_retry_after: true, +//! } +//! ``` + +use crate::classify::{classify_for_middleware, ClassifiedErrorKind}; +use crate::middleware::{ + MiddlewareContext, MiddlewareOutcome, PostProcessOutcome, TransportMiddleware, +}; +use gunbc_ir::transport::{RateLimitAlgorithm, RateLimitConfig, TransportRequest, TransportResponse}; +use std::collections::HashMap; +use std::sync::Mutex; +use std::time::{Duration, Instant}; + +/// Token bucket rate limiter state. +#[derive(Debug)] +struct TokenBucket { + /// Current number of available tokens. + tokens: f64, + /// Last time tokens were refilled. + last_refill: Instant, + /// Maximum tokens (burst capacity). + max_tokens: f64, + /// Tokens added per second. + refill_rate: f64, + /// Explicit pause until a specific time (for Retry-After handling). + /// If set and in the future, all requests are blocked. + pause_until: Option, +} + +impl TokenBucket { + fn new(max_burst: u32, sustained_per_minute: u32) -> Self { + // Ensure we don't divide by zero in try_acquire + // If sustained_per_minute is 0, use a very small rate (1 per hour) + let safe_rate = if sustained_per_minute == 0 { + 1.0 / 3600.0 // 1 per hour as minimum + } else { + sustained_per_minute as f64 / 60.0 + }; + + // max_burst of 0 means no requests allowed - use at least 1 + let safe_max = if max_burst == 0 { 1.0 } else { max_burst as f64 }; + + Self { + tokens: safe_max, + last_refill: Instant::now(), + max_tokens: safe_max, + refill_rate: safe_rate, + pause_until: None, + } + } + + /// Try to acquire a token. Returns wait time if rate limited. + fn try_acquire(&mut self) -> Result<(), Duration> { + let now = Instant::now(); + + // Check explicit pause first (from Retry-After) + if let Some(until) = self.pause_until { + if now < until { + return Err(until.duration_since(now)); + } + // Pause expired, clear it + self.pause_until = None; + } + + self.refill(); + + if self.tokens >= 1.0 { + self.tokens -= 1.0; + Ok(()) + } else { + // Calculate time to wait for next token + let tokens_needed = 1.0 - self.tokens; + let wait_secs = tokens_needed / self.refill_rate; + Err(Duration::from_secs_f64(wait_secs)) + } + } + + /// Refill tokens based on elapsed time. + fn refill(&mut self) { + let now = Instant::now(); + let elapsed = now.duration_since(self.last_refill); + let new_tokens = elapsed.as_secs_f64() * self.refill_rate; + self.tokens = (self.tokens + new_tokens).min(self.max_tokens); + self.last_refill = now; + } + + /// Current headroom as fraction (0.0 = exhausted, 1.0 = full). + fn headroom(&self) -> f64 { + // If paused, headroom is 0 + if let Some(until) = self.pause_until { + if Instant::now() < until { + return 0.0; + } + } + self.tokens / self.max_tokens + } + + /// Force a wait until a specific time (for Retry-After handling). + fn force_wait_until(&mut self, until: Instant) { + let now = Instant::now(); + if until > now { + // Set explicit pause - this is the only safe way to enforce wait + self.pause_until = Some(until); + // Also drain tokens to prevent any requests after pause expires + self.tokens = 0.0; + } + } +} + +/// Sliding window rate limiter state. +#[derive(Debug)] +struct SlidingWindow { + /// Request timestamps within the current window. + requests: Vec, + /// Window duration. + window: Duration, + /// Maximum requests per window. + max_requests: u32, + /// Explicit pause until a specific time (for Retry-After handling). + pause_until: Option, +} + +impl SlidingWindow { + fn new(sustained_per_minute: u32) -> Self { + Self { + requests: Vec::new(), + window: Duration::from_secs(60), + max_requests: sustained_per_minute, + pause_until: None, + } + } + + /// Try to record a request. Returns wait time if rate limited. + fn try_acquire(&mut self) -> Result<(), Duration> { + let now = Instant::now(); + + // Check explicit pause first (from Retry-After) + if let Some(until) = self.pause_until { + if now < until { + return Err(until.duration_since(now)); + } + // Pause expired, clear it + self.pause_until = None; + } + + self.prune(now); + + if (self.requests.len() as u32) < self.max_requests { + self.requests.push(now); + Ok(()) + } else { + // Calculate time until oldest request falls out of window + if let Some(oldest) = self.requests.first() { + // Use checked_duration_since to avoid panics on future timestamps + if let Some(oldest_age) = now.checked_duration_since(*oldest) { + if oldest_age < self.window { + let wait = self.window - oldest_age; + return Err(wait); + } + } else { + // oldest is in the future (shouldn't happen, but be safe) + return Err(Duration::from_secs(1)); + } + } + // Should not happen if prune worked correctly + Err(Duration::from_secs(1)) + } + } + + /// Remove requests older than the window. + fn prune(&mut self, now: Instant) { + // Use saturating_sub to avoid underflow when now < window + let cutoff = now.checked_sub(self.window).unwrap_or(now); + self.requests.retain(|t| *t > cutoff); + } + + /// Current headroom as fraction. + fn headroom(&self) -> f64 { + // If paused, headroom is 0 + if let Some(until) = self.pause_until { + if Instant::now() < until { + return 0.0; + } + } + let used = self.requests.len() as f64; + let max = self.max_requests as f64; + (max - used) / max + } + + /// Force a wait until a specific time (for Retry-After handling). + fn force_wait_until(&mut self, until: Instant) { + let now = Instant::now(); + if until > now { + // Use explicit pause instead of manipulating timestamps + // This avoids underflow panics when retry-after > window + self.pause_until = Some(until); + // Clear requests to be safe + self.requests.clear(); + } + } +} + +/// Rate limiter that supports both algorithms. +#[derive(Debug)] +enum RateLimiter { + TokenBucket(TokenBucket), + SlidingWindow(SlidingWindow), +} + +impl RateLimiter { + fn new(config: &RateLimitConfig) -> Self { + match config.algorithm { + RateLimitAlgorithm::TokenBucket => { + RateLimiter::TokenBucket(TokenBucket::new(config.max_burst, config.sustained_per_minute)) + } + RateLimitAlgorithm::SlidingWindow => { + RateLimiter::SlidingWindow(SlidingWindow::new(config.sustained_per_minute)) + } + } + } + + fn try_acquire(&mut self) -> Result<(), Duration> { + match self { + RateLimiter::TokenBucket(tb) => tb.try_acquire(), + RateLimiter::SlidingWindow(sw) => sw.try_acquire(), + } + } + + fn headroom(&self) -> f64 { + match self { + RateLimiter::TokenBucket(tb) => tb.headroom(), + RateLimiter::SlidingWindow(sw) => sw.headroom(), + } + } + + fn force_wait_until(&mut self, until: Instant) { + match self { + RateLimiter::TokenBucket(tb) => tb.force_wait_until(until), + RateLimiter::SlidingWindow(sw) => sw.force_wait_until(until), + } + } +} + +/// Shared rate limit state across requests. +#[derive(Debug, Default)] +pub struct RateLimitState { + /// Rate limiters keyed by scope. + limiters: Mutex>, +} + +impl RateLimitState { + pub fn new() -> Self { + Self::default() + } + + /// Get or create a rate limiter for a scope. + fn get_or_create(&self, scope: &str, config: &RateLimitConfig) -> Result<(), Duration> { + let mut limiters = self.limiters.lock().unwrap(); + let limiter = limiters + .entry(scope.to_string()) + .or_insert_with(|| RateLimiter::new(config)); + limiter.try_acquire() + } + + /// Get headroom for a scope. + fn headroom(&self, scope: &str) -> Option { + let limiters = self.limiters.lock().unwrap(); + limiters.get(scope).map(|l| l.headroom()) + } + + /// Apply Retry-After from a 429 response. + fn apply_retry_after(&self, scope: &str, retry_after_ms: u64) { + let mut limiters = self.limiters.lock().unwrap(); + if let Some(limiter) = limiters.get_mut(scope) { + let until = Instant::now() + Duration::from_millis(retry_after_ms); + limiter.force_wait_until(until); + } + } +} + +/// Rate limit middleware. +pub struct RateLimitMiddleware { + config: RateLimitConfig, + state: RateLimitState, +} + +impl RateLimitMiddleware { + pub fn new(config: RateLimitConfig) -> Self { + Self { + config, + state: RateLimitState::new(), + } + } + + /// Create with shared state (for use in pipeline). + pub fn with_shared_state(config: RateLimitConfig, state: RateLimitState) -> Self { + Self { config, state } + } +} + +impl TransportMiddleware for RateLimitMiddleware { + fn pre_request( + &self, + request: TransportRequest, + _ctx: &mut MiddlewareContext, + ) -> MiddlewareOutcome { + // Maximum retry attempts for rate limit acquisition. + // Prevents infinite loops while handling thundering herd scenarios where + // multiple threads wake up simultaneously after sleeping. + const MAX_ACQUIRE_ATTEMPTS: u32 = 10; + let mut attempts = 0; + + loop { + attempts += 1; + + match self.state.get_or_create(&self.config.scope_key, &self.config) { + Ok(()) => { + // Request allowed, record headroom for metrics + if let Some(headroom) = self.state.headroom(&self.config.scope_key) { + // Metrics sink would be called here via ctx.shared_state + let _ = headroom; + } + return MiddlewareOutcome::Continue(request); + } + Err(wait_duration) => { + if attempts >= MAX_ACQUIRE_ATTEMPTS { + // Too many failed attempts - abort rather than loop forever + return MiddlewareOutcome::Abort(gunbc_exec::ExecError::new(format!( + "Rate limit exhausted for scope '{}' after {} attempts", + self.config.scope_key, attempts + ))); + } + + // Rate limited - wait synchronously then retry + // In a real async implementation, this would be async sleep + std::thread::sleep(wait_duration); + // Loop continues to retry acquisition + } + } + } + } + + fn post_response( + &self, + _request: &TransportRequest, + response: TransportResponse, + ctx: &mut MiddlewareContext, + ) -> PostProcessOutcome { + // Check for 429 and apply Retry-After + if self.config.honor_retry_after { + if let Some(classified) = classify_for_middleware(&response, ctx.config.response_classification.as_ref()) { + if classified.kind == ClassifiedErrorKind::RateLimit { + if let Some(retry_after_ms) = classified.retry_after_ms { + self.state.apply_retry_after(&self.config.scope_key, retry_after_ms); + } + } + } + } + + PostProcessOutcome::Complete(response) + } + + fn name(&self) -> &'static str { + "rate_limit" + } +} + +#[cfg(test)] +mod tests { + use super::*; + use gunbc_ir::transport::{LocalRequest, TransportMiddlewareConfig}; + use std::sync::Arc; + use crate::middleware::SharedMiddlewareState; + + fn test_config(max_burst: u32, per_minute: u32) -> RateLimitConfig { + RateLimitConfig { + scope_key: "test".to_string(), + algorithm: RateLimitAlgorithm::TokenBucket, + max_burst, + sustained_per_minute: per_minute, + honor_retry_after: true, + } + } + + #[test] + fn token_bucket_allows_burst() { + let mut bucket = TokenBucket::new(5, 60); + + // Should allow 5 requests immediately (burst) + for _ in 0..5 { + assert!(bucket.try_acquire().is_ok()); + } + + // 6th request should be rate limited + assert!(bucket.try_acquire().is_err()); + } + + #[test] + fn token_bucket_refills_over_time() { + let mut bucket = TokenBucket::new(2, 120); // 2 tokens/sec + + // Use all tokens + assert!(bucket.try_acquire().is_ok()); + assert!(bucket.try_acquire().is_ok()); + assert!(bucket.try_acquire().is_err()); + + // Wait for refill + std::thread::sleep(Duration::from_millis(600)); + + // Should have ~1 token now + assert!(bucket.try_acquire().is_ok()); + } + + #[test] + fn sliding_window_limits_requests() { + let mut window = SlidingWindow::new(3); // 3 per minute + + assert!(window.try_acquire().is_ok()); + assert!(window.try_acquire().is_ok()); + assert!(window.try_acquire().is_ok()); + assert!(window.try_acquire().is_err()); + } + + #[test] + fn rate_limit_state_tracks_multiple_scopes() { + let state = RateLimitState::new(); + + let config1 = RateLimitConfig { + scope_key: "scope1".to_string(), + algorithm: RateLimitAlgorithm::TokenBucket, + max_burst: 2, + sustained_per_minute: 60, + honor_retry_after: false, + }; + + let config2 = RateLimitConfig { + scope_key: "scope2".to_string(), + algorithm: RateLimitAlgorithm::TokenBucket, + max_burst: 3, + sustained_per_minute: 60, + honor_retry_after: false, + }; + + // Each scope has its own bucket + assert!(state.get_or_create("scope1", &config1).is_ok()); + assert!(state.get_or_create("scope1", &config1).is_ok()); + assert!(state.get_or_create("scope1", &config1).is_err()); // scope1 exhausted + + // scope2 still has capacity + assert!(state.get_or_create("scope2", &config2).is_ok()); + assert!(state.get_or_create("scope2", &config2).is_ok()); + assert!(state.get_or_create("scope2", &config2).is_ok()); + assert!(state.get_or_create("scope2", &config2).is_err()); // scope2 exhausted + } + + #[test] + fn middleware_allows_request_within_limit() { + let config = test_config(10, 600); + let mw = RateLimitMiddleware::new(config); + + let mw_config = Arc::new(TransportMiddlewareConfig::default()); + let shared = Arc::new(SharedMiddlewareState::new()); + let mut ctx = MiddlewareContext::new("test.op", false, true, mw_config, shared); + + let request = TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}), + }); + + let outcome = mw.pre_request(request, &mut ctx); + assert!(matches!(outcome, MiddlewareOutcome::Continue(_))); + } + + #[test] + fn headroom_decreases_with_requests() { + let state = RateLimitState::new(); + let config = RateLimitConfig { + scope_key: "test".to_string(), + algorithm: RateLimitAlgorithm::TokenBucket, + max_burst: 4, + sustained_per_minute: 60, + honor_retry_after: false, + }; + + assert!(state.get_or_create("test", &config).is_ok()); + let h1 = state.headroom("test").unwrap(); + + assert!(state.get_or_create("test", &config).is_ok()); + let h2 = state.headroom("test").unwrap(); + + assert!(h2 < h1, "headroom should decrease: {} < {}", h2, h1); + } + + #[test] + fn apply_retry_after_blocks_requests() { + let state = RateLimitState::new(); + let config = RateLimitConfig { + scope_key: "test".to_string(), + algorithm: RateLimitAlgorithm::TokenBucket, + max_burst: 10, + sustained_per_minute: 60, + honor_retry_after: true, + }; + + // Prime the limiter + assert!(state.get_or_create("test", &config).is_ok()); + + // Apply a retry-after + state.apply_retry_after("test", 100); // 100ms + + // Should be rate limited now + let result = state.get_or_create("test", &config); + assert!(result.is_err()); + } +} diff --git a/lib/transport/src/retry.rs b/lib/transport/src/retry.rs new file mode 100644 index 00000000000..4fc6d54a7e7 --- /dev/null +++ b/lib/transport/src/retry.rs @@ -0,0 +1,672 @@ +//! Retry middleware with exponential/jittered backoff and circuit breaker. +//! +//! Provides automatic retry for transient failures (5xx, network errors, rate limits) +//! with configurable backoff strategies and circuit breaker protection. +//! +//! # Safety +//! +//! Only operations marked as `idempotent` or `readonly` are automatically retried. +//! Non-idempotent operations require explicit retry configuration. +//! +//! # Configuration +//! +//! ```ignore +//! RetryConfig { +//! max_attempts: 4, +//! base_delay_ms: 100, +//! max_delay_ms: 2000, +//! backoff: RetryBackoff::ExponentialJitter, +//! retry_statuses: vec![429, 500, 502, 503, 504], +//! retry_network_errors: true, +//! require_idempotent_or_readonly: true, +//! circuit_breaker: Some(CircuitBreakerConfig { ... }), +//! } +//! ``` + +use crate::classify::{classify_for_middleware, classify_transport_error}; +use crate::middleware::{ + MiddlewareContext, MiddlewareOutcome, PostProcessOutcome, TransportMiddleware, +}; +use gunbc_exec::ExecError; +use gunbc_ir::transport::{ + CircuitBreakerConfig, RetryBackoff, RetryConfig, TransportRequest, TransportResponse, +}; +use std::sync::Mutex; +use std::time::{Duration, Instant}; + +/// Circuit breaker state machine. +#[derive(Debug, Clone)] +pub enum CircuitState { + /// Normal operation, requests flow through. + Closed { + /// Consecutive failure count. + failure_count: u32, + }, + /// Requests blocked, waiting for reset timeout. + Open { + /// Time when circuit opened. + opened_at: Instant, + /// Time to wait before half-open. + reset_timeout: Duration, + }, + /// Testing recovery, limited requests allowed. + HalfOpen { + /// Successful probes in half-open state. + success_count: u32, + /// Failed probes in half-open state. + failure_count: u32, + /// Maximum probe requests before deciding. + max_probes: u32, + /// Number of requests currently in-flight (to prevent thundering herd). + in_flight: u32, + }, +} + +impl CircuitState { + fn new() -> Self { + CircuitState::Closed { failure_count: 0 } + } +} + +/// Circuit breaker protecting against cascading failures. +#[derive(Debug)] +pub struct CircuitBreaker { + config: CircuitBreakerConfig, + state: Mutex, +} + +impl CircuitBreaker { + pub fn new(config: CircuitBreakerConfig) -> Self { + Self { + config, + state: Mutex::new(CircuitState::new()), + } + } + + /// Check if request should be allowed. + /// + /// When transitioning from Open to HalfOpen, only ONE request is allowed at a time + /// to prevent thundering herd (all waiting requests rushing through at once). + pub fn should_allow(&self) -> bool { + let mut state = self.state.lock().unwrap(); + match &*state { + CircuitState::Closed { .. } => true, + CircuitState::Open { + opened_at, + reset_timeout, + } => { + // Check if it's time to half-open + if opened_at.elapsed() >= *reset_timeout { + // Transition to HalfOpen with one probe in flight + *state = CircuitState::HalfOpen { + success_count: 0, + failure_count: 0, + max_probes: self.config.half_open_max_requests, + in_flight: 1, // This request is the first probe + }; + true + } else { + false + } + } + CircuitState::HalfOpen { + success_count, + failure_count, + max_probes, + in_flight, + } => { + // Prevent thundering herd: only allow one request at a time in half-open. + // Wait for current probe to complete before allowing more. + if *in_flight > 0 { + return false; + } + + // Allow probe if under limit + if (success_count + failure_count) < *max_probes { + *state = CircuitState::HalfOpen { + success_count: *success_count, + failure_count: *failure_count, + max_probes: *max_probes, + in_flight: in_flight + 1, + }; + true + } else { + false + } + } + } + } + + /// Record a successful request. + pub fn record_success(&self) { + let mut state = self.state.lock().unwrap(); + match &*state { + CircuitState::Closed { .. } => { + // Reset failure count on success + *state = CircuitState::Closed { failure_count: 0 }; + } + CircuitState::HalfOpen { + success_count, + max_probes, + .. + } => { + let new_success = success_count + 1; + if new_success >= *max_probes { + // Recovered, close circuit + *state = CircuitState::Closed { failure_count: 0 }; + } else { + // Probe succeeded, allow next probe (in_flight = 0) + *state = CircuitState::HalfOpen { + success_count: new_success, + failure_count: 0, + max_probes: *max_probes, + in_flight: 0, + }; + } + } + CircuitState::Open { .. } => { + // Should not happen - success while open + } + } + } + + /// Record a failed request. + pub fn record_failure(&self) { + let mut state = self.state.lock().unwrap(); + match &*state { + CircuitState::Closed { failure_count } => { + let new_count = failure_count + 1; + if new_count >= self.config.failure_threshold { + // Trip the circuit + *state = CircuitState::Open { + opened_at: Instant::now(), + reset_timeout: Duration::from_millis(self.config.reset_timeout_ms), + }; + } else { + *state = CircuitState::Closed { + failure_count: new_count, + }; + } + } + CircuitState::HalfOpen { .. } => { + // Probe failed, reopen circuit + *state = CircuitState::Open { + opened_at: Instant::now(), + reset_timeout: Duration::from_millis(self.config.reset_timeout_ms), + }; + } + CircuitState::Open { .. } => { + // Already open, nothing to do + } + } + } + + /// Check if circuit is open. + pub fn is_open(&self) -> bool { + matches!(&*self.state.lock().unwrap(), CircuitState::Open { .. }) + } +} + +/// Calculate backoff delay for retry attempt. +fn calculate_backoff(config: &RetryConfig, attempt: u32) -> Duration { + let base = config.base_delay_ms as f64; + let max = config.max_delay_ms as f64; + + let delay_ms = match config.backoff { + RetryBackoff::Fixed => base, + RetryBackoff::Exponential => { + // 2^(attempt-1) * base, capped at max + let factor = 2_f64.powi((attempt - 1) as i32); + (base * factor).min(max) + } + RetryBackoff::ExponentialJitter => { + // Exponential with random jitter ±25% + let factor = 2_f64.powi((attempt - 1) as i32); + let base_delay = (base * factor).min(max); + // True random jitter in range [0.75, 1.25] to avoid thundering herd + use std::collections::hash_map::RandomState; + use std::hash::{BuildHasher, Hasher}; + let mut hasher = RandomState::new().build_hasher(); + hasher.write_u64(std::time::Instant::now().elapsed().as_nanos() as u64); + let random_bits = hasher.finish(); + let jitter_factor = 0.75 + ((random_bits % 500) as f64 / 1000.0); + base_delay * jitter_factor + } + }; + + Duration::from_millis(delay_ms as u64) +} + +/// Check if a status code should be retried. +fn is_retryable_status(config: &RetryConfig, status: u16) -> bool { + config.retry_statuses.contains(&status) +} + +/// Retry middleware. +pub struct RetryMiddleware { + config: RetryConfig, + circuit_breaker: Option, +} + +impl RetryMiddleware { + pub fn new(config: RetryConfig) -> Self { + let circuit_breaker = config + .circuit_breaker + .as_ref() + .map(|cb| CircuitBreaker::new(cb.clone())); + Self { + config, + circuit_breaker, + } + } + + /// Check if operation can be retried. + fn can_retry(&self, ctx: &MiddlewareContext) -> bool { + if self.config.require_idempotent_or_readonly { + ctx.retry_safe() + } else { + true + } + } + + /// Check if response indicates a retryable error. + fn should_retry_response( + &self, + response: &TransportResponse, + ctx: &MiddlewareContext, + ) -> Option { + // Check if we've exceeded max attempts + if ctx.attempt >= self.config.max_attempts { + return None; + } + + // Check if operation allows retry + if !self.can_retry(ctx) { + return None; + } + + // Classify the response + if let Some(classified) = + classify_for_middleware(response, ctx.config.response_classification.as_ref()) + { + // Check if error kind is retryable + if !classified.retryable() { + return None; + } + + // Check if specific status is in retry list + if let Some(status) = classified.status { + if !is_retryable_status(&self.config, status) { + return None; + } + } + + // Retryable + return Some( + classified + .message + .unwrap_or_else(|| format!("{:?}", classified.kind)), + ); + } + + None + } + + /// Check if error indicates a retryable condition. + fn should_retry_error(&self, _error: &ExecError, ctx: &MiddlewareContext) -> Option { + // Check limits + if ctx.attempt >= self.config.max_attempts { + return None; + } + + if !self.can_retry(ctx) { + return None; + } + + // Network errors are retryable if configured + if self.config.retry_network_errors { + Some("network error".to_string()) + } else { + None + } + } +} + +impl TransportMiddleware for RetryMiddleware { + fn pre_request( + &self, + request: TransportRequest, + _ctx: &mut MiddlewareContext, + ) -> MiddlewareOutcome { + // Check circuit breaker + if let Some(cb) = &self.circuit_breaker { + if !cb.should_allow() { + return MiddlewareOutcome::Abort(ExecError::new( + "Circuit breaker is open - too many recent failures", + )); + } + } + + MiddlewareOutcome::Continue(request) + } + + fn post_response( + &self, + _request: &TransportRequest, + response: TransportResponse, + ctx: &mut MiddlewareContext, + ) -> PostProcessOutcome { + // Check if this was a success (for circuit breaker) + // Only Server and Network errors should trip the circuit breaker. + // Client errors (4xx including 404) are typically the caller's fault, + // not a sign of service degradation. + let classified = + classify_for_middleware(&response, ctx.config.response_classification.as_ref()); + + // For HTTP transports, use classification; for non-HTTP, use is_success + let is_circuit_breaker_failure = if let Some(c) = &classified { + // HTTP failure: Server or Network errors trip the circuit breaker + matches!( + c.kind, + crate::classify::ClassifiedErrorKind::Server + | crate::classify::ClassifiedErrorKind::Network + ) + } else { + // Non-HTTP transport (File, Shell, Tcp, Local): use is_success + // If the transport reports failure, count it for circuit breaker + !crate::classify::is_success(&response) + }; + + if let Some(cb) = &self.circuit_breaker { + if is_circuit_breaker_failure { + cb.record_failure(); + } else { + cb.record_success(); + } + } + + // Check if we should retry + if let Some(reason) = self.should_retry_response(&response, ctx) { + let delay = calculate_backoff(&self.config, ctx.attempt); + return PostProcessOutcome::Retry { + delay_ms: delay.as_millis() as u64, + reason, + }; + } + + PostProcessOutcome::Complete(response) + } + + fn on_error( + &self, + _request: &TransportRequest, + error: ExecError, + ctx: &mut MiddlewareContext, + ) -> PostProcessOutcome { + // Don't process synthetic pipeline cleanup errors - just pass through + // without affecting circuit breaker state or attempting retry + let error_msg = error.to_string(); + if error_msg.contains("pipeline cleanup") { + return PostProcessOutcome::Abort(error); + } + + // Record failure for circuit breaker + if let Some(cb) = &self.circuit_breaker { + cb.record_failure(); + } + + // Classify the error + let classified = classify_transport_error(&error_msg); + let is_retryable = classified.retryable(); + + // Check if we should retry + if is_retryable { + if let Some(reason) = self.should_retry_error(&error, ctx) { + let delay = calculate_backoff(&self.config, ctx.attempt); + return PostProcessOutcome::Retry { + delay_ms: delay.as_millis() as u64, + reason, + }; + } + } + + PostProcessOutcome::Abort(error) + } + + fn name(&self) -> &'static str { + "retry" + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::middleware::SharedMiddlewareState; + use gunbc_ir::transport::{LocalRequest, RestResponse, TransportMiddlewareConfig}; + use std::sync::Arc; + + fn basic_retry_config() -> RetryConfig { + RetryConfig { + max_attempts: 3, + base_delay_ms: 100, + max_delay_ms: 1000, + backoff: RetryBackoff::Exponential, + retry_statuses: vec![429, 500, 502, 503, 504], + retry_network_errors: true, + require_idempotent_or_readonly: true, + circuit_breaker: None, + } + } + + fn cb_config() -> CircuitBreakerConfig { + CircuitBreakerConfig { + failure_threshold: 3, + reset_timeout_ms: 100, + half_open_max_requests: 2, + } + } + + #[test] + fn backoff_fixed_returns_constant() { + let config = RetryConfig { + base_delay_ms: 100, + max_delay_ms: 1000, + backoff: RetryBackoff::Fixed, + ..basic_retry_config() + }; + + assert_eq!(calculate_backoff(&config, 1), Duration::from_millis(100)); + assert_eq!(calculate_backoff(&config, 2), Duration::from_millis(100)); + assert_eq!(calculate_backoff(&config, 3), Duration::from_millis(100)); + } + + #[test] + fn backoff_exponential_doubles() { + let config = RetryConfig { + base_delay_ms: 100, + max_delay_ms: 10000, + backoff: RetryBackoff::Exponential, + ..basic_retry_config() + }; + + assert_eq!(calculate_backoff(&config, 1), Duration::from_millis(100)); + assert_eq!(calculate_backoff(&config, 2), Duration::from_millis(200)); + assert_eq!(calculate_backoff(&config, 3), Duration::from_millis(400)); + assert_eq!(calculate_backoff(&config, 4), Duration::from_millis(800)); + } + + #[test] + fn backoff_exponential_caps_at_max() { + let config = RetryConfig { + base_delay_ms: 100, + max_delay_ms: 300, + backoff: RetryBackoff::Exponential, + ..basic_retry_config() + }; + + assert_eq!(calculate_backoff(&config, 1), Duration::from_millis(100)); + assert_eq!(calculate_backoff(&config, 2), Duration::from_millis(200)); + assert_eq!(calculate_backoff(&config, 3), Duration::from_millis(300)); // capped + assert_eq!(calculate_backoff(&config, 4), Duration::from_millis(300)); // capped + } + + #[test] + fn circuit_breaker_starts_closed() { + let cb = CircuitBreaker::new(cb_config()); + assert!(cb.should_allow()); + assert!(!cb.is_open()); + } + + #[test] + fn circuit_breaker_opens_on_threshold() { + let cb = CircuitBreaker::new(cb_config()); + + // Record failures up to threshold + cb.record_failure(); + assert!(cb.should_allow()); + cb.record_failure(); + assert!(cb.should_allow()); + cb.record_failure(); + + // Should now be open + assert!(cb.is_open()); + assert!(!cb.should_allow()); + } + + #[test] + fn circuit_breaker_success_resets_count() { + let cb = CircuitBreaker::new(cb_config()); + + cb.record_failure(); + cb.record_failure(); + cb.record_success(); // Reset + cb.record_failure(); + cb.record_failure(); + + // Should still be closed (count reset) + assert!(cb.should_allow()); + } + + #[test] + fn circuit_breaker_half_open_after_timeout() { + let mut config = cb_config(); + config.reset_timeout_ms = 10; // Very short for test + let cb = CircuitBreaker::new(config); + + // Open the circuit + cb.record_failure(); + cb.record_failure(); + cb.record_failure(); + assert!(cb.is_open()); + + // Wait for timeout + std::thread::sleep(Duration::from_millis(20)); + + // Should allow (half-open) + assert!(cb.should_allow()); + } + + #[test] + fn retry_middleware_allows_idempotent() { + let mw = RetryMiddleware::new(basic_retry_config()); + let mw_config = Arc::new(TransportMiddlewareConfig::default()); + let shared = Arc::new(SharedMiddlewareState::new()); + + // Idempotent operation + let ctx = MiddlewareContext::new("test", true, false, mw_config.clone(), shared.clone()); + assert!(mw.can_retry(&ctx)); + + // Readonly operation + let ctx = MiddlewareContext::new("test", false, true, mw_config.clone(), shared.clone()); + assert!(mw.can_retry(&ctx)); + + // Non-idempotent, non-readonly + let ctx = MiddlewareContext::new("test", false, false, mw_config.clone(), shared.clone()); + assert!(!mw.can_retry(&ctx)); + } + + #[test] + fn retry_middleware_respects_max_attempts() { + let mw = RetryMiddleware::new(basic_retry_config()); + let mw_config = Arc::new(TransportMiddlewareConfig::default()); + let shared = Arc::new(SharedMiddlewareState::new()); + let mut ctx = MiddlewareContext::new("test", true, false, mw_config, shared); + + let response = TransportResponse::Rest(RestResponse::new( + 500, + serde_json::json!({"error": "server error"}), + )); + + // Attempt 1 - should retry + ctx.attempt = 1; + let outcome = mw.post_response(&TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}), + }), response.clone(), &mut ctx); + assert!(matches!(outcome, PostProcessOutcome::Retry { .. })); + + // Attempt 2 - should retry + ctx.attempt = 2; + let outcome = mw.post_response(&TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}), + }), response.clone(), &mut ctx); + assert!(matches!(outcome, PostProcessOutcome::Retry { .. })); + + // Attempt 3 (max) - should not retry + ctx.attempt = 3; + let outcome = mw.post_response(&TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}), + }), response.clone(), &mut ctx); + assert!(matches!(outcome, PostProcessOutcome::Complete(_))); + } + + #[test] + fn retry_middleware_checks_status_code() { + let mw = RetryMiddleware::new(basic_retry_config()); + let mw_config = Arc::new(TransportMiddlewareConfig::default()); + let shared = Arc::new(SharedMiddlewareState::new()); + let mut ctx = MiddlewareContext::new("test", true, false, mw_config, shared); + ctx.attempt = 1; + + // 500 should retry + let response = TransportResponse::Rest(RestResponse::new(500, serde_json::json!({}))); + let outcome = mw.post_response(&TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}), + }), response, &mut ctx); + assert!(matches!(outcome, PostProcessOutcome::Retry { .. })); + + // 400 should not retry (not in retry_statuses) + let response = TransportResponse::Rest(RestResponse::new(400, serde_json::json!({}))); + let outcome = mw.post_response(&TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}), + }), response, &mut ctx); + assert!(matches!(outcome, PostProcessOutcome::Complete(_))); + } + + #[test] + fn retry_middleware_blocks_when_circuit_open() { + let config = RetryConfig { + circuit_breaker: Some(cb_config()), + ..basic_retry_config() + }; + let mw = RetryMiddleware::new(config); + let mw_config = Arc::new(TransportMiddlewareConfig::default()); + let shared = Arc::new(SharedMiddlewareState::new()); + let mut ctx = MiddlewareContext::new("test", true, false, mw_config, shared); + + // Record failures to open circuit + let response = TransportResponse::Rest(RestResponse::new(500, serde_json::json!({}))); + for _ in 0..3 { + ctx.attempt = 1; + mw.post_response(&TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}), + }), response.clone(), &mut ctx); + } + + // Now circuit should be open + let request = TransportRequest::Local(LocalRequest { + inputs: serde_json::json!({}), + }); + let outcome = mw.pre_request(request, &mut ctx); + assert!(matches!(outcome, MiddlewareOutcome::Abort(_))); + } +} diff --git a/lib/transport/src/transport_types.rs b/lib/transport/src/transport_types.rs new file mode 100644 index 00000000000..d9cc5f2496e --- /dev/null +++ b/lib/transport/src/transport_types.rs @@ -0,0 +1,200 @@ +//! Transport foundation types for classifying and configuring transport operations. +//! +//! These types provide semantic annotations that guide middleware behavior, +//! test generation, and runtime optimization. + +/// Transport class - the fundamental protocol/mechanism. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum TransportClass { + /// REST over HTTP/HTTPS. + Rest, + /// Raw HTTP without REST semantics. + Http, + /// Shell command execution. + Shell, + /// File system operations. + File, + /// Raw TCP sockets. + Tcp, + /// gRPC (future). + Grpc, + /// Streaming connections (SSE, WebSocket). + Stream, + /// Pub/sub messaging (future). + Pubsub, + /// Custom/plugin transport. + Custom, +} + +impl TransportClass { + /// Whether this transport class supports connection pooling. + pub fn supports_pooling(&self) -> bool { + matches!(self, Self::Rest | Self::Http | Self::Grpc | Self::Tcp) + } + + /// Whether this transport class is inherently streaming. + pub fn is_streaming(&self) -> bool { + matches!(self, Self::Stream | Self::Pubsub) + } + + /// Whether this transport class has request/response semantics. + pub fn is_request_response(&self) -> bool { + matches!(self, Self::Rest | Self::Http | Self::Grpc | Self::Shell | Self::File) + } +} + +/// Capabilities of a transport implementation. +#[derive(Debug, Clone, PartialEq, Eq, Default)] +pub struct TransportCapabilities { + /// Supports connection pooling for reuse. + pub connection_pooling: bool, + /// Operations are safe to retry on failure. + pub retry_safe: bool, + /// Supports streaming responses. + pub streaming: bool, + /// Supports bidirectional communication. + pub bidirectional: bool, + /// Supports multiplexing multiple requests over one connection. + pub multiplexing: bool, + /// Has built-in compression support. + pub compression: bool, +} + +/// Endpoint-level behavior hints. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum EndpointBehavior { + /// Endpoint is rate-limited; respect limits. + RateLimited, + /// Response can be cached. + Cacheable, + /// Operation is idempotent; safe to retry. + Idempotent, + /// Response is paginated; may need multiple requests. + Paginated, + /// Endpoint requires authentication. + Authenticated, + /// Endpoint supports conditional requests (ETag/If-Modified-Since). + Conditional, + /// Endpoint is deprecated; warn on use. + Deprecated, + /// Endpoint has known latency; adjust timeouts. + HighLatency, +} + +/// Operation-level behavioral properties. +/// +/// These flags guide middleware decisions about retry, caching, and observability. +#[derive(Debug, Clone, PartialEq, Eq, Default)] +pub struct OperationBehavior { + /// Operation only reads state; no side effects. + pub readonly: bool, + /// Operation can be safely retried without side effects. + pub idempotent: bool, + /// Operation has no external dependencies (for testing). + pub hermetic: bool, + /// Expected behaviors for this endpoint. + pub behaviors: Vec, + /// Custom retry timeout override (ms). + pub timeout_ms: Option, + /// Maximum retry attempts override. + pub max_retries: Option, + /// Known failure modes for this operation. + pub failure_modes: Vec, +} + +impl OperationBehavior { + /// Create a new readonly operation. + pub fn readonly() -> Self { + Self { + readonly: true, + idempotent: true, // readonly implies idempotent + ..Default::default() + } + } + + /// Create a new idempotent operation. + pub fn idempotent() -> Self { + Self { + idempotent: true, + ..Default::default() + } + } + + /// Create a new hermetic operation (for testing). + pub fn hermetic() -> Self { + Self { + hermetic: true, + ..Default::default() + } + } + + /// Whether this operation is safe to retry. + pub fn retry_safe(&self) -> bool { + self.readonly || self.idempotent + } + + /// Check if a specific behavior is declared. + pub fn has_behavior(&self, behavior: EndpointBehavior) -> bool { + self.behaviors.contains(&behavior) + } +} + +/// Known failure modes for an operation. +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum FailureMode { + /// Operation may time out under load. + Timeout, + /// Operation may be rate limited. + RateLimited, + /// Operation may fail due to auth issues. + AuthenticationRequired, + /// Operation may fail due to missing resource. + NotFound, + /// Operation may fail due to conflict. + Conflict, + /// Operation may fail due to validation errors. + ValidationError, + /// Operation may fail due to quota exhaustion. + QuotaExceeded, + /// Custom failure mode with description. + Custom(String), +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn transport_class_capabilities() { + assert!(TransportClass::Rest.supports_pooling()); + assert!(TransportClass::Rest.is_request_response()); + assert!(!TransportClass::Rest.is_streaming()); + + assert!(TransportClass::Stream.is_streaming()); + assert!(!TransportClass::Shell.supports_pooling()); + } + + #[test] + fn operation_behavior_retry_safe() { + let readonly = OperationBehavior::readonly(); + assert!(readonly.retry_safe()); + + let idempotent = OperationBehavior::idempotent(); + assert!(idempotent.retry_safe()); + + let default = OperationBehavior::default(); + assert!(!default.retry_safe()); + } + + #[test] + fn operation_behavior_has_behavior() { + let behavior = OperationBehavior { + behaviors: vec![EndpointBehavior::RateLimited, EndpointBehavior::Cacheable], + ..Default::default() + }; + + assert!(behavior.has_behavior(EndpointBehavior::RateLimited)); + assert!(behavior.has_behavior(EndpointBehavior::Cacheable)); + assert!(!behavior.has_behavior(EndpointBehavior::Deprecated)); + } +} diff --git a/tasks.md b/tasks.md index 434dc97b0d3..46aa5d484fe 100644 --- a/tasks.md +++ b/tasks.md @@ -539,9 +539,10 @@ spec.rs # Service operation specs | C17 | RT96 | **Kill `propagate_to_param_sources`.** Fix boundary detection. Param source nodes auto-fed. | `propagate_to_param_sources` deleted. One port per input. | M | | C18 | — | **Executor dead code.** Delete `looks_effectful_without_kind()`. Delete unwired credential expiry plumbing. | Dead code deleted. `cargo clippy` clean. | S | | C19 | RT83, RT4b | **Restore passthrough enforcement + runtime fail-closed diagnostics.** After C4+C5+C7 wire dag_util branches, required outputs with no input must return `ExecError` (not `Skipped`) and emit clear diagnostics for missing declared passthroughs (RT4b). | `resolve.rs` returns `ExecError` for required missing outputs. Missing passthrough ports are diagnosable (no silent fallback). CI clean (no unwired branches). | S | -| C20 | RT59, RT63 | **CLI generator: profile, mode, subcommand support.** Expose `available_profiles` in `CompileOutput`. Template generates `--profile` enum flag, `--mode ensure\|verify`, subcommand dispatch for multi-func modules. `KEY=VALUE` arg parsing for infra-style tools. Unblocks Worker A. | Generated CLI for `pipelines/sdlc.dag` accepts `--profile`. Generated CLI for multi-func modules has subcommands. | L | +| ~~C20~~ | ~~RT59, RT63~~ | ~~**CLI generator: profile, mode, subcommand support.** Expose `available_profiles` in `CompileOutput`. Template generates `--profile` enum flag, `--mode ensure\|verify`, subcommand dispatch for multi-func modules. Unblocks Worker A.~~ **Done** | ~~Generated CLI for `pipelines/sdlc.dag` accepts `--profile`. Generated CLI for multi-func modules has subcommands.~~ | ~~L~~ | +| C21 | — | **CLI generator: KEY=VALUE and multi-value flag support.** For `Map` params, generate `KEY=VALUE` parser (e.g., `--input project_id=my-project`). For `List` params, generate accumulator flags (`--target A --target B`). Required for A5 (infra.rs elimination). | `gunbc-infra --input project_id=foo` parses to map. `--target A --target B` parses to list. | M | -**Chain**: C1 → C3; C2; C10 (RT4a/c) → C4 → C5 → C6; C7; C8; C9; C10a → (RF-INV1 or RF-INV2 gate) → C11 → C14 → C15 → C19; C12; C13; C16; C17; C18; C20 (early, unblocks A) +**Chain**: C1 → C3; C2; C10 (RT4a/c) → C4 → C5 → C6; C7; C8; C9; C10a → (RF-INV1 or RF-INV2 gate) → C11 → C14 → C15 → C19; C12; C13; C16; C17; C18; C20 (early, unblocks A); C21 (unblocks A5) --- @@ -551,6 +552,7 @@ The following tasks are done. Kept for postmortem/audit reference. | ID | What | Status | |----|------|--------| +| C20 | CLI generator: profile, mode, subcommand support (RT59, RT63) | Done | | RT1 | Credential wiring (`auth_input` → `res:credential`) | Done | | RT2 | Execute node fail-closed when `auth_scheme` declared | Done | | RT3 | File transport: EXISTS, CREATE_DIR, DELETE, APPEND, GLOB | Done | @@ -961,11 +963,24 @@ the sibling repos useful (failure modes, edge cases, rate limits, prerequisites) ### Principle -Make the transport layer production-ready: rate limiting, retry with backoff, -response classification, credential middleware, enriched virtual backends. -All runtime infrastructure that Lane 7 (SDLC Production) needs. +Two-phase approach following the Protobuf/gRPC pattern: -**Rust work in `lib/transport/`, `core/test/`, `core/ir/src/transport/`.** +**Phase 1 (TL-0:10): Target SDK** — Build language-specific OS mechanisms (token +buckets, mutexes, retry loops, credential caches). This is the "runtime library" +that handles thread sleep, atomic counters, TCP sockets. Domain-agnostic. + +**Phase 2 (TL-11:15): Domain Modeling** — Move domain policy (rate limit budgets, +retry rules, error shapes) from Rust code into `.dag` service definitions. The +compiler generates **configuration code** that links to the Target SDK. + +The distinction: +- **Domain policy (What)** → `.dag` — "GitHub rate limit is 5000/hour" +- **OS mechanisms (How)** → Target SDK — "How do I atomically decrement a counter?" + +**Design doc**: `docs/design/transport-primitives.md` + +**Phase 1 work**: `lib/transport/`, `core/test/`, `core/ir/src/transport/` +**Phase 2 work**: `core/daglang/`, `dsl/services/` ### Why this lane exists @@ -975,9 +990,9 @@ varies per service), automatic retry with jittered backoff, response classification (status code → typed error before field extraction), credential refresh/caching, and enriched virtual backends for hermetic testing. -The virtual transport backend (`test_backend.rs`) handles File ops and basic -Shell but HTTP stubs are unimplemented. Lane 7 (SDLC Production) needs all -of these for real execution and for fidelity-tiered hermetic tests. +Phase 1 builds these mechanisms as a domain-agnostic Target SDK. Phase 2 +moves the domain-specific configuration (rate limits, error shapes) into +`.dag` files so the compiler can emit configuration for any target language. ### File Territory @@ -988,11 +1003,15 @@ of these for real execution and for fidelity-tiered hermetic tests. **No overlap with Red Team** (which works in `core/daglang/`, `core/codegen/`, `core/exec/`, `gunbc-dag/src/`). -### Queue +### Phase 1: Target SDK (TL-0:10) + +Build the domain-agnostic OS-level middleware. This code doesn't know "GitHub" +or "GCP" — it only knows "I was configured with budget=5000, window=3600". | # | ID | Task | Size | Status | Deps | |---|-----|------|------|--------|------| -| 1 | TL-1 | **Rate limit middleware.** `lib/transport/src/rate_limit.rs`. Token bucket + sliding window implementations. Per-endpoint rate tracking. `RateLimitConfig` from IR transport metadata. Automatic 429/retry-after handling. Shared rate state across concurrent requests. Tests: burst exhaustion, recovery, concurrent access. | L | Pending | — | +| 0 | TL-0 | **Transport foundation types.** `lib/transport/src/transport_types.rs`. `TransportClass` enum (Rest, Shell, File, Grpc, Stream, Pubsub, Custom). `TransportCapabilities` struct (connection_pooling, retry_safe, streaming, etc.). `EndpointBehavior` enum (RateLimited, Cacheable, Idempotent, Paginated, etc.). `OperationBehavior` struct (readonly, idempotent, hermetic flags + retry/timeout config). Exported from `lib.rs`. | M | **Done** | — | +| 1 | TL-1 | **Rate limit middleware.** `lib/transport/src/rate_limit.rs`. Token bucket + sliding window implementations. Per-endpoint rate tracking. `RateLimitConfig` from IR transport metadata. Automatic 429/retry-after handling. Shared rate state across concurrent requests. Tests: burst exhaustion, recovery, concurrent access. | L | Pending | TL-0 | | 2 | TL-2 | **Retry middleware with backoff.** `lib/transport/src/retry.rs`. Exponential + jittered backoff. Configurable per-operation via `RetryPolicy` from IR. Idempotency-aware: only auto-retry operations marked `idempotent` or `readonly`. Transient error classification (5xx, network timeout, rate limit). Circuit breaker for persistent failures. Tests: retry sequences, backoff timing, circuit breaker state machine. | L | Pending | TL-1 | | 3 | TL-3 | **Response classification.** `lib/transport/src/classify.rs`. HTTP status code → typed error mapping. Per-provider error shape parsing (GitHub `{ message, documentation_url }`, GCP `{ error: { code, message, status } }`, Anthropic `{ type, error: { type, message } }`). Classification hierarchy: auth error > rate limit > client error > server error > network error. Feeds error into retry decision (TL-2). Tests: per-provider error shapes, unknown shapes, malformed responses. | M | Pending | — | | 4 | TL-4 | **Credential middleware.** `lib/transport/src/credential.rs`. Token caching with TTL-aware refresh. Multi-provider credential resolution (OAuth2 bearer, GCP WIF, API key). Automatic credential injection into requests based on `AuthScheme` from service config. Proactive refresh at 80% TTL. Thread-safe credential store. Tests: token refresh, concurrent access, expired token detection. | L | Pending | TL-3 | @@ -1005,6 +1024,28 @@ of these for real execution and for fidelity-tiered hermetic tests. **Chain**: TL-1 → TL-2 → TL-9; TL-3 → TL-4, TL-7; TL-5; TL-6; TL-10 (early, no deps); TL-8 → TL-9 +### Phase 2: Transport Domain Modeling (after TL-9) + +**Design doc**: `docs/design/transport-primitives.md` + +**Principle**: Phase 1 (TL-0:10) builds the **Target SDK** — OS-level mechanisms (token buckets, mutexes, sockets) that are language-specific. Phase 2 moves **domain data** (rate limit budgets, retry policies, error shapes) into `.dag` where it belongs. + +The distinction: +- **Domain policy (What)** → `.dag` — "GitHub rate limit is 5000/hour" is an external fact about a service +- **OS mechanisms (How)** → Target SDK — "How do I atomically decrement a counter?" is Rust/Go/Python specific + +This follows the Protobuf/gRPC pattern: the compiler generates **configuration code** that links to a Target SDK, not line-by-line OS implementations. + +| # | ID | Task | Size | Status | Deps | +|---|-----|------|------|--------|------| +| 11 | TL-11 | **DSL syntax for transport blocks.** Add `rate_limit {}`, `retry {}`, `error_shape {}`, `credential {}` blocks to grammar. Parse budget expressions (`5000 per hour`). Typecheck scope bindings (`uses rate_limit: core`). | L | Pending | TL-9 | +| 12 | TL-12 | **Lower transport blocks to IR.** Lower DSL transport blocks to existing `TransportMiddlewareConfig` IR (TL-10). Rate limit budgets become `RateLimitConfig`. Retry policies become `RetryConfig`. Rust runtime still interprets. | M | Pending | TL-10, TL-11 | +| 13 | TL-13 | **Domain data migration.** Move hardcoded rate limits from Rust to `dsl/services/*.dag`. GitHub 5000/hour core, 30/min search. GCP quotas. Anthropic limits. Delete provider-specific branches from `classify.rs`. Service definitions become source of truth. | M | Pending | TL-12 | +| 14 | TL-14 | **Multi-target emit.** Emit transport configuration per target language. Rust emits code linking to existing Target SDK. Go/Python stubs for future. Generated code calls `RateLimitMiddleware::new(config)` — doesn't reimplement token bucket. | XL | Pending | TL-13 | +| 15 | TL-15 | **Substrate cleanup.** `lib/transport/` becomes pure Target SDK (no domain knowledge). Delete `GITHUB_CORE_LIMIT` constants, `host.contains("github.com")` branches. All domain facts live in `.dag`. | L | Pending | TL-14 | + +**Chain (Phase 2)**: TL-11 → TL-12 → TL-13 → TL-14 → TL-15 + --- ## Lane 6: Service Layer Completion (after ED lane + Red Team + Lane 4)