mirror of
https://github.com/block/buzz.git
synced 2026-08-18 06:50:31 +02:00
## What Implements Phases 2 and 4a of the Usage v2 plan (plan events `d0268cd0`/`0e95b035`), extending the archive backend to emit, transport, archive, and aggregate both cache categories and billing identity fail-closed. ### P2 — emission, transport, archive **Tri-state accumulators** (`Unseen`/`Exact`/`Unknown`) for cache-read and cache-write in `buzz-agent` turn and session state. Absent field = Unknown (never zero) through the full pipeline. No `unwrap_or(0)` on the cache path. Both cache folds are gated on usage-bearing responses (same gate as the total-state and identity folds) — a response with no usage at all must not poison either accumulator. **Overflow-aware input token parsing and accumulation** — closed end-to-end from parse through wire to ACP: - `sum_usage()` returns `SumUsageResult` (`Exact(u64)` | `Overflow`) — checked arithmetic, never clamps. `anthropic_input_tokens()` returns `Option<SumUsageResult>` since it sums three fields (`input_tokens + cache_read_input_tokens + cache_creation_input_tokens`) that can collectively overflow. Single-field callers (`prompt_tokens`, `completion_tokens`, etc.) convert via `.into_exact()` — their single-field sums cannot overflow. - `LlmResponse.input_tokens_overflowed: bool` propagates the parse-layer signal into the run loop. When set, `input_tokens` is `None` (clamped value discarded), the context-gate baseline (`last_request_input_tokens`) is frozen at its prior reading, and `turn_input_tokens` is poisoned to `TurnIOState::Poisoned` before any emission — including mid-turn `emit_usage_update` calls. A dedicated enum on `LlmResponse.input_tokens` would ripple into ~20 existing test assertions on `r.input_tokens == Some(...)`; the bool flag confines the change to the two call sites that check it. - `TurnIOState` (`Unseen`/`Exact`/`Poisoned`) for input and output: per-round fold uses `checked_add`; overflow poisons permanently at turn and session level, no healing. Absence does not poison (pass-2-cleared contract unchanged). Wire emission omits `accumulatedInputTokens`/`accumulatedOutputTokens` when poisoned — never null, never `u64::MAX`. ACP treats absent = publisher-poisoned: `delta_reliable: false`, null turn fields, null cumulative for that category; session cumulative stays unknown for all subsequent turns once poisoned. **Conditional wire emission** for `accumulatedCachedInputTokens` and new `accumulatedCacheWriteTokens`: fields are omitted when the cumulative is Unseen or Unknown. ACP `_goose/unstable/session/update` contract documented next to the payload with tests for all absence/zero variants. **`PricingIdentity` stamping (publisher-side)**: - `pricing_authority()`: canonical parsed-URL endpoint comparison against the official allowlist — HTTPS only, exact allowlisted host (lookalike-safe), default port (omitted or explicit :443), required API base path, rejects userinfo/query/fragment/path-prefix lookalikes. - Model: the actually-requested `request_model` after mesh/auto resolution (not `effective_model_str`). - Turn discipline: identity retained only while ALL usage in the current turn carries one identical proven identity; any mismatch, unproven-usage-bearing response, or unpaired cumulative snapshot poisons to absent; a later matching notification does not heal a mixed turn. **ACP `UsageTracker` identity fold**: per-in-flight-turn tri-state identity accumulator replacing last-update-wins. Any absent identity on a token-advancing notification or exact mismatch poisons to absent; poison survives later updates; reset in `begin_turn()`/`take()`; reset also when a request fails (baseline cleared so preflight gate cannot stay frozen sub-threshold on retries). **M3 migration**: adds `turn_cache_write_tokens`, `cumulative_cache_write_tokens`, `pricing_authority`, `pricing_model`, `pricing_cache_class` to `agent_metric_index`. Additive, idempotent, guarded per-column by marker. M2 migration also guarded per-column (turn and cumulative cache-read columns checked and added independently; marker commits only after both are present). Fresh-DB schema includes all columns. **First-turn baselines**: `seed_zero_baseline` seeds `last_input: Some(0)`, `last_output: Some(0)`, `last_cached_input: Some(0)`, `last_cache_write: Some(0)`, and `last_total: Some(0)` — all have the known-zero-at-spawn argument. Absent fields from incoming snapshots still produce unknown (tri-state unchanged). Sessions buzz-acp did not spawn (no seed) remain fail-closed on turn one. **`ReportedUsage` TS mirror**: `cacheReadTokens`, `cacheWriteTokens`, `freshInputTokens` added to `tauriArchive.ts` as `UsageField` members, field-for-field with the Rust struct. ### P4a — aggregation layer **Extended S-1 ladder** to cache-read and cache-write via the same `ladder_token` path as the existing token fields. **`freshInputTokens` derivation**: checked arithmetic, fail-closed — absent cache fields produce Unknown (not zero), overflow and `cacheRead+cacheWrite > input` both produce `incomplete: true`. Aggregated as a `UsageField`. **D6 comparator**: `sort_value()` = provider total when known, else `input+output` when both known, else `None` (unknown-last). Replaces the prior total-only comparator for both agent-level and model-level sort. Ships a pinned test vector that the TS render layer (P5) must match. ## Test coverage - `buzz-agent`: 440 lib + 15 integration (golden_transcripts) — includes 13 new `cache_total_state_tests`; 14 new `turn_io_state_tests`; 3 new `sum_usage_*` tests (exact single-field, exact two-field, overflow signals correctly); 3 new `parse_anthropic_*` tests (overflow flag set + value cleared, normal sum no flag, absent usage no flag); end-to-end golden transcript drives real subprocess with Anthropic-shaped `input_tokens: u64::MAX, cache_read: 1` response and asserts `accumulatedInputTokens` absent from the emitted `usage_update` — no logic duplication; 3 wire pin tests; 4 `fold_pricing_identity_*` tests; `pricing_authority()` explicit-:443 acceptance - `buzz-acp`: 700 tests (691 lib + 9 integration) — 4 new usage tests (absent input → unreliable+null; absent output → unreliable+null; goose-shaped both present unchanged; poison mid-session); 3 ACP behavior tests; 7 pool lifecycle tests - Desktop (Rust): 2259+ tests — 14 new P4a pinned tests; 2 M3 round-trip tests; 1 serde key-shape test; 2 M2 partial-schema migration tests; first-turn cache round-trip test ## Related PRs - P1 NIP-AM spec: [#4632](https://github.com/block/buzz/pull/4632) - P3 pricing table: [#4629](https://github.com/block/buzz/pull/4629) - UI (P5): [#4001](https://github.com/block/buzz/pull/4001) --------- Signed-off-by: Will Pfleger <pfleger.will@gmail.com> Signed-off-by: Duncan <dcfd242e557282d7a1e2cf2e6877522682f1e5c6156dc92ca7d90eaedd3b0f95@buzz.block.builderlab.xyz> Co-authored-by: npub1mn7jgtj4w2pd0g0zeuhxsa6jy6p0rewxz4kujt98my82ahfmp72sxjexk7 <dcfd242e557282d7a1e2cf2e6877522682f1e5c6156dc92ca7d90eaedd3b0f95@buzz.block.builderlab.xyz>
1022 lines
39 KiB
Rust
1022 lines
39 KiB
Rust
#![forbid(unsafe_code)]
|
|
mod agent;
|
|
pub mod auth;
|
|
mod builtin;
|
|
pub mod catalog;
|
|
pub mod config;
|
|
mod handoff;
|
|
mod hints;
|
|
mod llm;
|
|
mod mcp;
|
|
pub mod types;
|
|
mod wire;
|
|
|
|
pub use catalog::{discover_databricks_models, ModelEntry, DATABRICKS_V2_KNOWN_MODELS};
|
|
pub use config::Provider;
|
|
pub use types::AgentError;
|
|
|
|
/// Environment keys the Windows Git Bash resolver may inspect. `spawn_one()`
|
|
/// forwards every key in this list into its otherwise-cleared MCP child; Doctor
|
|
/// uses the same contract so a ready agent can always start its shell tool.
|
|
#[cfg(windows)]
|
|
pub const WINDOWS_SHELL_RESOLUTION_ENV: &[&str] = &[
|
|
"PATH",
|
|
"BUZZ_SHELL",
|
|
"GIT_BASH",
|
|
"SystemRoot",
|
|
"ProgramFiles",
|
|
"ProgramFiles(x86)",
|
|
"LOCALAPPDATA",
|
|
];
|
|
|
|
use std::collections::HashMap;
|
|
use std::path::Path;
|
|
use std::sync::Arc;
|
|
|
|
use serde_json::{json, Value};
|
|
use tokio::io::BufReader;
|
|
use tokio::sync::{mpsc, watch, Mutex};
|
|
|
|
use crate::agent::RunCtx;
|
|
use crate::config::{Config, MAX_SYSTEM_PROMPT_BYTES, PROTOCOL_VERSION};
|
|
use crate::hints::SkillEntry;
|
|
use crate::llm::Llm;
|
|
use crate::mcp::McpRegistry;
|
|
use crate::types::{ContentBlock, HistoryItem};
|
|
use crate::wire::{
|
|
classify, goose_session_update, Inbound, InitializeParams, SessionCancelParams,
|
|
SessionNewParams, SessionPromptParams, SessionSetModelParams, SessionSteerParams, WireMsg,
|
|
WireSender, INVALID_PARAMS, METHOD_NOT_FOUND, PARSE_ERROR,
|
|
};
|
|
|
|
struct App {
|
|
cfg: Config,
|
|
llm: Arc<Llm>,
|
|
sessions: Mutex<HashMap<String, Session>>,
|
|
/// Cached model catalog for Databricks providers. Populated lazily on the
|
|
/// first successful `session/new` discovery call. Failed discovery is never
|
|
/// cached: static-token authentication errors reject session creation, while
|
|
/// OAuth authentication and non-auth errors use the configured model for that
|
|
/// response and retry on the next session.
|
|
models_cache: tokio::sync::OnceCell<Vec<ModelEntry>>,
|
|
}
|
|
|
|
struct Session {
|
|
id: String,
|
|
mcp: Arc<McpRegistry>,
|
|
/// Skills discovered at session creation; used by the built-in `load_skill` tool.
|
|
skills: Vec<SkillEntry>,
|
|
history: Vec<HistoryItem>,
|
|
cancel_tx: watch::Sender<bool>,
|
|
busy: bool,
|
|
/// Run id of the in-flight prompt, set when a prompt starts and cleared
|
|
/// when it ends. `None` means no active run — a steer request targeting
|
|
/// this session is rejected. Steer-capable clients learn this value from
|
|
/// the `params.update._meta.goose.activeRunId` field on `session/update`.
|
|
active_run_id: Option<String>,
|
|
/// Sender for mid-turn steer messages. Created fresh per prompt (like
|
|
/// `cancel_tx`); the running prompt loop holds the matching receiver and
|
|
/// drains queued steers at round boundaries. `None` when no prompt is in
|
|
/// flight.
|
|
steer_tx: Option<mpsc::UnboundedSender<Vec<ContentBlock>>>,
|
|
original_task: Option<String>,
|
|
handoff_count: usize,
|
|
/// Cache-summed input tokens the provider reported for this session's most
|
|
/// recent request, or `None` before the first response (or after a handoff
|
|
/// resets the context). Drives the token-based handoff gate; see
|
|
/// [`RunCtx::should_handoff`].
|
|
last_request_input_tokens: Option<u64>,
|
|
/// History byte size when `last_request_input_tokens` was measured, paired
|
|
/// with it so the gate can account for history appended since.
|
|
last_request_history_bytes: Option<usize>,
|
|
effective_system_prompt: Arc<str>,
|
|
/// Per-session model override set by `session/set_model`. When `Some`,
|
|
/// overrides `App::cfg.model` for all LLM calls on this session. Persists
|
|
/// across `session/prompt` calls until changed.
|
|
effective_model: Option<String>,
|
|
/// Session-cumulative input tokens across all turns. Sent in the
|
|
/// `_goose/unstable/session/update` usage notification so buzz-acp's
|
|
/// `UsageTracker` can compute per-turn deltas symmetrically with goose.
|
|
/// `TurnIOState`: `Unseen` before any turn reports; `Exact(n)` while running;
|
|
/// `Poisoned` if any turn's sum overflowed — permanently poisons the session.
|
|
accumulated_input_tokens: crate::types::TurnIOState,
|
|
/// Session-cumulative output tokens across all turns.
|
|
/// Same `Unseen`/`Exact(n)`/`Poisoned` contract as `accumulated_input_tokens`.
|
|
accumulated_output_tokens: crate::types::TurnIOState,
|
|
/// Session-cumulative cache-served input tokens across all turns — a subset
|
|
/// of `accumulated_input_tokens`, not an addition to it. Tri-state:
|
|
///
|
|
/// - `Unseen`: no turn has ever reported this category.
|
|
/// - `Exact(n)`: every usage-bearing response in every turn reported this
|
|
/// category; `n` is the cumulative sum.
|
|
/// - `Unknown`: at least one usage-bearing response ever omitted the
|
|
/// category — permanently poisoned for this session.
|
|
accumulated_cached_input_tokens: crate::types::CacheTotalState,
|
|
/// Session-cumulative cache-written input tokens across all turns — also a
|
|
/// subset of `accumulated_input_tokens`, not an addition to it.
|
|
/// Same `Unseen`/`Exact`/`Unknown` tri-state contract as
|
|
/// `accumulated_cached_input_tokens`.
|
|
accumulated_cache_write_tokens: crate::types::CacheTotalState,
|
|
/// Session-cumulative total-token state across all turns.
|
|
///
|
|
/// Mirrors the per-turn `TurnTotalState` tri-state: starts `Unseen`,
|
|
/// becomes `Exact(n)` as turns with genuine provider totals complete,
|
|
/// transitions permanently to `Unknown` when any turn lacks a total or
|
|
/// when the cumulative would otherwise decrease. Only emitted in the
|
|
/// `usage_update` notification when `Exact`.
|
|
accumulated_total_state: crate::types::TurnTotalState,
|
|
}
|
|
|
|
fn die(msg: String) -> ! {
|
|
tracing::error!("{msg}");
|
|
std::process::exit(2);
|
|
}
|
|
|
|
pub fn run() -> Result<(), Box<dyn std::error::Error>> {
|
|
let args: Vec<String> = std::env::args().collect();
|
|
if matches!(args.get(1).map(String::as_str), Some("auth")) {
|
|
return tokio::runtime::Builder::new_multi_thread()
|
|
.enable_all()
|
|
.build()?
|
|
.block_on(auth_subcommand(&args[2..]));
|
|
}
|
|
tokio::runtime::Builder::new_multi_thread()
|
|
.enable_all()
|
|
.build()?
|
|
.block_on(async_main());
|
|
Ok(())
|
|
}
|
|
|
|
pub async fn authenticate_databricks(host: &str) -> Result<(), AgentError> {
|
|
auth::PkceOAuthTokenSource::new(llm::databricks_pkce_config(host))?
|
|
.interactive_login()
|
|
.await
|
|
}
|
|
|
|
/// `buzz-agent auth <provider>` — run the interactive auth flow for a
|
|
/// provider and persist the result, then exit. Today this supports Databricks
|
|
/// OAuth 2.0 PKCE. Reads `DATABRICKS_HOST` from env; needs a browser on the
|
|
/// machine.
|
|
async fn auth_subcommand(args: &[String]) -> Result<(), Box<dyn std::error::Error>> {
|
|
let provider = args.first().map(String::as_str);
|
|
match provider {
|
|
Some("databricks" | "databricks_v2" | "databricks-v2") => {
|
|
let host = std::env::var("DATABRICKS_HOST")
|
|
.map_err(|_| "auth databricks: DATABRICKS_HOST required")?;
|
|
authenticate_databricks(&host).await?;
|
|
eprintln!("Authenticated. Token cached under ~/.config/buzz-agent/oauth/databricks/.");
|
|
Ok(())
|
|
}
|
|
Some(other) => Err(format!("auth: unknown provider {other:?}").into()),
|
|
None => Err("auth: provider required (try: buzz-agent auth databricks)".into()),
|
|
}
|
|
}
|
|
|
|
async fn async_main() {
|
|
tracing_subscriber::fmt()
|
|
.with_writer(std::io::stderr)
|
|
.with_ansi(false)
|
|
.init();
|
|
let cfg = Config::from_env().unwrap_or_else(|e| die(e));
|
|
let llm = Arc::new(Llm::new(&cfg).unwrap_or_else(|e| die(e.to_string())));
|
|
let max_line = cfg.max_line_bytes;
|
|
let app = Arc::new(App {
|
|
cfg,
|
|
llm,
|
|
sessions: Mutex::new(HashMap::new()),
|
|
models_cache: tokio::sync::OnceCell::new(),
|
|
});
|
|
let (wire_tx, wire_rx) = mpsc::channel::<WireMsg>(64);
|
|
let writer = tokio::spawn(wire::writer_task(wire_rx));
|
|
if let Err(e) = read_loop(
|
|
BufReader::new(tokio::io::stdin()),
|
|
app.clone(),
|
|
wire_tx,
|
|
max_line,
|
|
)
|
|
.await
|
|
{
|
|
tracing::error!("io: reader: {e}");
|
|
}
|
|
for session in app.sessions.lock().await.values() {
|
|
let _ = session.cancel_tx.send(true);
|
|
}
|
|
let _ = writer.await;
|
|
}
|
|
|
|
async fn read_loop<R: tokio::io::AsyncBufRead + Unpin>(
|
|
mut stdin: R,
|
|
app: Arc<App>,
|
|
wire_tx: WireSender,
|
|
max_line: usize,
|
|
) -> std::io::Result<()> {
|
|
while let Some(line) = wire::read_bounded_line(&mut stdin, max_line).await? {
|
|
if line.trim().is_empty() {
|
|
continue;
|
|
}
|
|
match serde_json::from_str::<Value>(&line) {
|
|
Ok(msg) => dispatch(&app, msg, &wire_tx).await,
|
|
Err(e) => {
|
|
wire::send(
|
|
&wire_tx,
|
|
wire::err(Value::Null, PARSE_ERROR, &format!("jsonrpc: parse: {e}")),
|
|
)
|
|
.await;
|
|
}
|
|
}
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
async fn dispatch(app: &Arc<App>, msg: Value, wire_tx: &WireSender) {
|
|
match classify(&msg) {
|
|
Inbound::Request { id, method, params } => {
|
|
handle_request(app, id, method, params, wire_tx).await
|
|
}
|
|
Inbound::Notification { method, params } => handle_notification(app, &method, params).await,
|
|
Inbound::Ignored => {}
|
|
Inbound::Invalid { id, code, message } => {
|
|
wire::send(wire_tx, wire::err(id, code, &message)).await
|
|
}
|
|
}
|
|
}
|
|
|
|
async fn handle_request(
|
|
app: &Arc<App>,
|
|
id: Value,
|
|
method: String,
|
|
params: Value,
|
|
wire_tx: &WireSender,
|
|
) {
|
|
match method.as_str() {
|
|
"initialize" => initialize(id, params, wire_tx).await,
|
|
"session/new" => {
|
|
let app = app.clone();
|
|
let wire_tx = wire_tx.clone();
|
|
tokio::spawn(async move { session_new(&app, id, params, &wire_tx).await });
|
|
}
|
|
"session/prompt" => spawn_prompt(app.clone(), id, params, wire_tx.clone()),
|
|
"session/set_model" => {
|
|
set_model_session(app, id, params, wire_tx).await;
|
|
}
|
|
"session/cancel" => {
|
|
cancel_session(app, params).await;
|
|
wire::send(wire_tx, wire::ok(id, Value::Null)).await;
|
|
}
|
|
// goose-compatible non-standard extension: inject user input into the
|
|
// currently active prompt without starting a new one. Mirrors goose's
|
|
// `_goose/unstable/session/steer` wire contract so a single client-side
|
|
// delivery path serves both agents.
|
|
"_goose/unstable/session/steer" => {
|
|
steer_session(app, id, params, wire_tx).await;
|
|
}
|
|
_ => {
|
|
wire::send(
|
|
wire_tx,
|
|
wire::err(
|
|
id,
|
|
METHOD_NOT_FOUND,
|
|
&format!("jsonrpc: method not found: {method}"),
|
|
),
|
|
)
|
|
.await
|
|
}
|
|
}
|
|
}
|
|
|
|
async fn handle_notification(app: &Arc<App>, method: &str, params: Value) {
|
|
if method == "session/cancel" {
|
|
cancel_session(app, params).await;
|
|
}
|
|
}
|
|
|
|
async fn initialize(id: Value, params: Value, wire_tx: &WireSender) {
|
|
let p: InitializeParams = match decode(params, "initialize") {
|
|
Ok(p) => p,
|
|
Err(m) => return reject(wire_tx, id, INVALID_PARAMS, &m).await,
|
|
};
|
|
// Honest negotiation: respond with the minimum of what the client
|
|
// requested and what we support.
|
|
// NOTE: gating `[Base]` injection on `protocol_version < 2` is a deliberate
|
|
// temporary measure — we are squatting on ACP v2 ahead of the upstream ACP
|
|
// RFD. Revisit when that RFD merges; otherwise a genuine upstream-v2 agent
|
|
// would silently lose `[Base]`.
|
|
let negotiated_version = p.protocol_version.min(PROTOCOL_VERSION);
|
|
wire::send(
|
|
wire_tx,
|
|
wire::ok(
|
|
id,
|
|
json!({
|
|
"protocolVersion": negotiated_version,
|
|
"agentCapabilities": {
|
|
"loadSession": false,
|
|
"promptCapabilities": { "image": false, "audio": false, "embeddedContext": false },
|
|
"mcpCapabilities": { "http": false, "sse": false },
|
|
},
|
|
"agentInfo": { "name": "buzz-agent", "version": env!("CARGO_PKG_VERSION") },
|
|
}),
|
|
),
|
|
)
|
|
.await;
|
|
}
|
|
|
|
/// Resolve the Databricks model catalog for one `session/new` call.
|
|
///
|
|
/// Tries to use a previously-cached successful discovery result. If the cache is empty,
|
|
/// runs `discover` and — on success — populates the cache for future calls. On failure
|
|
/// the error is returned and the cell is intentionally left empty so the next session retries.
|
|
///
|
|
/// Extracted from `session_new` so that tests can drive this path with an injected
|
|
/// discovery future without requiring a full `App` / transport stack.
|
|
async fn resolve_models_catalog(
|
|
cache: &tokio::sync::OnceCell<Vec<ModelEntry>>,
|
|
discover: impl std::future::Future<Output = Result<Vec<ModelEntry>, AgentError>>,
|
|
) -> Result<Vec<ModelEntry>, AgentError> {
|
|
cache.get_or_try_init(|| discover).await.cloned()
|
|
}
|
|
|
|
/// Return the configured model as a one-entry catalog for this response.
|
|
///
|
|
/// This value is never written to `models_cache`; failed discovery must be retried by
|
|
/// the next session rather than pinning degraded state for the process lifetime.
|
|
fn configured_model_fallback(model: &str) -> Vec<ModelEntry> {
|
|
let model = model.trim().to_string();
|
|
vec![ModelEntry {
|
|
id: model.clone(),
|
|
name: model,
|
|
}]
|
|
}
|
|
|
|
async fn session_new(app: &Arc<App>, id: Value, params: Value, wire_tx: &WireSender) {
|
|
let p: SessionNewParams = match decode(params, "session/new") {
|
|
Ok(p) => p,
|
|
Err(m) => return reject(wire_tx, id, INVALID_PARAMS, &m).await,
|
|
};
|
|
if p.cwd.is_empty() || !Path::new(&p.cwd).is_absolute() {
|
|
return reject(
|
|
wire_tx,
|
|
id,
|
|
INVALID_PARAMS,
|
|
"session/new: cwd must be an absolute path",
|
|
)
|
|
.await;
|
|
}
|
|
// Check cap without holding lock across MCP spawn (which may be slow).
|
|
{
|
|
let sessions = app.sessions.lock().await;
|
|
if sessions.len() >= app.cfg.max_sessions {
|
|
return reject(
|
|
wire_tx,
|
|
id,
|
|
INVALID_PARAMS,
|
|
"session/new: max sessions reached",
|
|
)
|
|
.await;
|
|
}
|
|
}
|
|
let (hints_text, skills) = if app.cfg.hints_enabled {
|
|
hints::build_hints_section(std::path::Path::new(&p.cwd))
|
|
} else {
|
|
(String::new(), Vec::new())
|
|
};
|
|
let effective_system_prompt: Arc<str> = {
|
|
// When the harness provides a systemPrompt (base_prompt + persona), use
|
|
// it as the primary content and suppress the default. The default is only
|
|
// a fallback for legacy harnesses that don't send systemPrompt.
|
|
let base = match p.system_prompt.as_deref() {
|
|
Some(client_prompt) if !client_prompt.trim().is_empty() => client_prompt.to_owned(),
|
|
_ => app.cfg.system_prompt.clone(),
|
|
};
|
|
let prompt = if hints_text.is_empty() {
|
|
base
|
|
} else {
|
|
format!("{base}\n\n{hints_text}")
|
|
};
|
|
// Reject combined prompts exceeding 512KB.
|
|
if prompt.len() > MAX_SYSTEM_PROMPT_BYTES {
|
|
return reject(
|
|
wire_tx,
|
|
id,
|
|
INVALID_PARAMS,
|
|
&format!(
|
|
"session/new: combined system prompt exceeds {}KB limit ({} bytes)",
|
|
MAX_SYSTEM_PROMPT_BYTES / 1024,
|
|
prompt.len()
|
|
),
|
|
)
|
|
.await;
|
|
}
|
|
Arc::from(prompt)
|
|
};
|
|
// Resolve the model catalog before spawning MCP servers or registering a
|
|
// session. A configured static credential cannot recover interactively, so
|
|
// its authentication failure rejects before allocation. OAuth authentication
|
|
// failures and other catalog failures use only the configured model for this
|
|
// response, without caching, so session/prompt can run the existing PKCE flow.
|
|
let available_models: Vec<Value> = {
|
|
use crate::config::Provider;
|
|
match app.cfg.provider {
|
|
Provider::Databricks | Provider::DatabricksV2 => {
|
|
let models = match resolve_models_catalog(
|
|
&app.models_cache,
|
|
discover_databricks_models(&app.cfg),
|
|
)
|
|
.await
|
|
{
|
|
Ok(models) => models,
|
|
Err(error @ AgentError::LlmAuth(_)) if !app.cfg.api_key.is_empty() => {
|
|
return reject(wire_tx, id, error.json_rpc_code(), &error.to_string())
|
|
.await;
|
|
}
|
|
Err(error @ AgentError::LlmAuth(_)) => {
|
|
tracing::warn!(
|
|
error = %error,
|
|
"Databricks OAuth model catalog unavailable; using configured model"
|
|
);
|
|
configured_model_fallback(&app.cfg.model)
|
|
}
|
|
Err(error) => {
|
|
tracing::warn!(
|
|
error = %error,
|
|
"Databricks model catalog unavailable; using configured model"
|
|
);
|
|
configured_model_fallback(&app.cfg.model)
|
|
}
|
|
};
|
|
models
|
|
.iter()
|
|
.map(|m| json!({ "modelId": m.id, "name": m.name }))
|
|
.collect()
|
|
}
|
|
_ => vec![json!({ "modelId": app.cfg.model, "name": app.cfg.model })],
|
|
}
|
|
};
|
|
|
|
let mcp = match McpRegistry::spawn_all(&app.cfg, &p.mcp_servers, &p.cwd).await {
|
|
Ok(m) => Arc::new(m),
|
|
Err(e) => return reject(wire_tx, id, e.json_rpc_code(), &e.to_string()).await,
|
|
};
|
|
let session_id = match session_token() {
|
|
Ok(t) => format!("ses_{t}"),
|
|
Err(e) => return reject(wire_tx, id, -32000, &e).await,
|
|
};
|
|
let (cancel_tx, _) = watch::channel(false);
|
|
let mut sessions = app.sessions.lock().await;
|
|
// Re-check cap (another session may have been created while we spawned MCP).
|
|
if sessions.len() >= app.cfg.max_sessions {
|
|
return reject(
|
|
wire_tx,
|
|
id,
|
|
INVALID_PARAMS,
|
|
"session/new: max sessions reached",
|
|
)
|
|
.await;
|
|
}
|
|
sessions.insert(
|
|
session_id.clone(),
|
|
Session {
|
|
id: session_id.clone(),
|
|
mcp,
|
|
skills,
|
|
history: Vec::new(),
|
|
cancel_tx,
|
|
busy: false,
|
|
active_run_id: None,
|
|
steer_tx: None,
|
|
original_task: None,
|
|
handoff_count: 0,
|
|
last_request_input_tokens: None,
|
|
last_request_history_bytes: None,
|
|
effective_system_prompt,
|
|
effective_model: None,
|
|
accumulated_input_tokens: crate::types::TurnIOState::Unseen,
|
|
accumulated_output_tokens: crate::types::TurnIOState::Unseen,
|
|
accumulated_cached_input_tokens: crate::types::CacheTotalState::Unseen,
|
|
accumulated_cache_write_tokens: crate::types::CacheTotalState::Unseen,
|
|
accumulated_total_state: crate::types::TurnTotalState::Unseen,
|
|
},
|
|
);
|
|
drop(sessions);
|
|
|
|
wire::send(
|
|
wire_tx,
|
|
wire::ok(
|
|
id,
|
|
json!({
|
|
"sessionId": session_id,
|
|
"models": {
|
|
"currentModelId": app.cfg.model,
|
|
"availableModels": available_models,
|
|
},
|
|
}),
|
|
),
|
|
)
|
|
.await;
|
|
}
|
|
|
|
fn decode<T: serde::de::DeserializeOwned>(params: Value, stage: &str) -> Result<T, String> {
|
|
serde_json::from_value(params).map_err(|e| format!("{stage}: {e}"))
|
|
}
|
|
|
|
async fn reject(wire_tx: &WireSender, id: Value, code: i32, message: &str) {
|
|
wire::send(wire_tx, wire::err(id, code, message)).await;
|
|
}
|
|
|
|
async fn cancel_session(app: &Arc<App>, params: Value) {
|
|
if let Ok(p) = serde_json::from_value::<SessionCancelParams>(params) {
|
|
if let Some(s) = app.sessions.lock().await.get(&p.session_id) {
|
|
let _ = s.cancel_tx.send(true);
|
|
}
|
|
}
|
|
}
|
|
|
|
/// Handle `session/set_model`: apply a per-session model override immediately.
|
|
///
|
|
/// Validation:
|
|
/// - Unknown `sessionId` → `invalid_params`.
|
|
/// - Empty `modelId` → `invalid_params`.
|
|
///
|
|
/// On success: stores `model_id` on the session and responds `{ sessionId, modelId }`.
|
|
/// The override is picked up by the next `session/prompt` call on this session.
|
|
async fn set_model_session(app: &Arc<App>, id: Value, params: Value, wire_tx: &WireSender) {
|
|
let p: SessionSetModelParams = match decode(params, "session/set_model") {
|
|
Ok(p) => p,
|
|
Err(m) => return reject(wire_tx, id, INVALID_PARAMS, &m).await,
|
|
};
|
|
if p.model_id.trim().is_empty() {
|
|
return reject(
|
|
wire_tx,
|
|
id,
|
|
INVALID_PARAMS,
|
|
"session/set_model: modelId must not be empty",
|
|
)
|
|
.await;
|
|
}
|
|
let mut sessions = app.sessions.lock().await;
|
|
let Some(s) = sessions.get_mut(&p.session_id) else {
|
|
return reject(
|
|
wire_tx,
|
|
id,
|
|
INVALID_PARAMS,
|
|
"session/set_model: unknown session",
|
|
)
|
|
.await;
|
|
};
|
|
s.effective_model = Some(p.model_id.clone());
|
|
tracing::info!(
|
|
session_id = %p.session_id,
|
|
model_id = %p.model_id,
|
|
"session/set_model: model overridden"
|
|
);
|
|
drop(sessions);
|
|
wire::send(
|
|
wire_tx,
|
|
wire::ok(
|
|
id,
|
|
json!({ "sessionId": p.session_id, "modelId": p.model_id }),
|
|
),
|
|
)
|
|
.await;
|
|
}
|
|
|
|
/// Handle `_goose/unstable/session/steer`: queue user input into the in-flight
|
|
/// prompt. Validation mirrors goose's `on_steer_session`:
|
|
/// - empty prompt → `invalid_params`
|
|
/// - no active run (no prompt in flight) → `invalid_params`
|
|
/// - `expectedRunId` mismatch → `invalid_params` (caller is steering a turn
|
|
/// that already ended or rotated; it must fall back to cancel+merge)
|
|
///
|
|
/// On success the message is queued for pickup at the next round boundary and
|
|
/// we reply `{ runId, messageId }`, then emit a `queuedSteer` session/update so
|
|
/// the client can correlate the accepted steer with its eventual pickup.
|
|
async fn steer_session(app: &Arc<App>, id: Value, params: Value, wire_tx: &WireSender) {
|
|
let p: SessionSteerParams = match decode(params, "_goose/unstable/session/steer") {
|
|
Ok(p) => p,
|
|
Err(m) => return reject(wire_tx, id, INVALID_PARAMS, &m).await,
|
|
};
|
|
if p.prompt.is_empty() {
|
|
return reject(
|
|
wire_tx,
|
|
id,
|
|
INVALID_PARAMS,
|
|
"steer: prompt must not be empty",
|
|
)
|
|
.await;
|
|
}
|
|
if p.expected_run_id.is_empty() {
|
|
return reject(
|
|
wire_tx,
|
|
id,
|
|
INVALID_PARAMS,
|
|
"steer: expectedRunId must not be empty",
|
|
)
|
|
.await;
|
|
}
|
|
let message_id = format!("steer_{}", session_token().unwrap_or_else(|_| "x".into()));
|
|
let run_id = {
|
|
let sessions = app.sessions.lock().await;
|
|
let Some(s) = sessions.get(&p.session_id) else {
|
|
return reject(wire_tx, id, INVALID_PARAMS, "steer: unknown session").await;
|
|
};
|
|
let Some(active) = s.active_run_id.as_deref() else {
|
|
return reject(wire_tx, id, INVALID_PARAMS, "steer: no active run to steer").await;
|
|
};
|
|
if active != p.expected_run_id {
|
|
return reject(
|
|
wire_tx,
|
|
id,
|
|
INVALID_PARAMS,
|
|
&format!(
|
|
"steer: expected active run id `{}` but found `{active}`",
|
|
p.expected_run_id
|
|
),
|
|
)
|
|
.await;
|
|
}
|
|
// A live run always has a steer_tx; if the channel is gone the run is
|
|
// tearing down — treat as no active run rather than queue into the void.
|
|
match &s.steer_tx {
|
|
Some(tx) if tx.send(p.prompt).is_ok() => active.to_owned(),
|
|
_ => return reject(wire_tx, id, INVALID_PARAMS, "steer: no active run to steer").await,
|
|
}
|
|
};
|
|
wire::send(
|
|
wire_tx,
|
|
wire::ok(id, json!({ "runId": run_id, "messageId": message_id })),
|
|
)
|
|
.await;
|
|
// Best-effort correlation hint for the client; mirrors goose's
|
|
// `send_queued_steer_update`. Not load-bearing for delivery.
|
|
wire::send(
|
|
wire_tx,
|
|
wire::session_update_with_goose_meta(
|
|
&p.session_id,
|
|
json!({ "sessionUpdate": "session_info_update" }),
|
|
json!({ "queuedSteer": { "messageId": message_id, "runId": run_id } }),
|
|
),
|
|
)
|
|
.await;
|
|
}
|
|
|
|
fn spawn_prompt(app: Arc<App>, id: Value, params: Value, wire_tx: WireSender) {
|
|
tokio::spawn(async move { run_prompt(app, id, params, wire_tx).await });
|
|
}
|
|
|
|
async fn run_prompt(app: Arc<App>, id: Value, params: Value, wire_tx: WireSender) {
|
|
let p: SessionPromptParams = match decode(params, "session/prompt") {
|
|
Ok(p) => p,
|
|
Err(m) => return reject(&wire_tx, id, INVALID_PARAMS, &m).await,
|
|
};
|
|
let (
|
|
sid,
|
|
mcp,
|
|
skills,
|
|
mut history,
|
|
mut original_task,
|
|
mut handoff_count,
|
|
mut last_request_input_tokens,
|
|
mut last_request_history_bytes,
|
|
mut cancel_rx,
|
|
effective_system_prompt,
|
|
effective_model_override,
|
|
run_id,
|
|
mut steer_rx,
|
|
usage_baseline,
|
|
) = match acquire_session(&app, &p.session_id).await {
|
|
Ok(v) => v,
|
|
Err(reason) => {
|
|
return reject(
|
|
&wire_tx,
|
|
id,
|
|
INVALID_PARAMS,
|
|
&format!("session/prompt: {reason}"),
|
|
)
|
|
.await
|
|
}
|
|
};
|
|
// Advertise the active run id so steer-capable clients can target this turn
|
|
// via `expectedRunId`. Mirrors goose's `send_active_run_update`.
|
|
wire::send(
|
|
&wire_tx,
|
|
wire::session_update_with_goose_meta(
|
|
&sid,
|
|
json!({ "sessionUpdate": "session_info_update" }),
|
|
json!({ "activeRunId": run_id }),
|
|
),
|
|
)
|
|
.await;
|
|
// Resolve effective model: session override wins over config default.
|
|
let effective_model_str = effective_model_override
|
|
.as_deref()
|
|
.unwrap_or(&app.cfg.model);
|
|
let mut turn_input_tokens: crate::types::TurnIOState = crate::types::TurnIOState::Unseen;
|
|
let mut turn_output_tokens: crate::types::TurnIOState = crate::types::TurnIOState::Unseen;
|
|
let mut turn_cached_input_tokens: crate::types::CacheTotalState =
|
|
crate::types::CacheTotalState::Unseen;
|
|
let mut turn_cache_write_tokens: crate::types::CacheTotalState =
|
|
crate::types::CacheTotalState::Unseen;
|
|
let mut turn_total_state = crate::types::TurnTotalState::Unseen;
|
|
// Per-turn billing identity accumulator — three-state:
|
|
// None = no usage-bearing response seen yet (initial)
|
|
// Some(Some(pi))= all usage-bearing responses carry the same proven identity
|
|
// Some(None) = poisoned (mixed identities, unproven response, etc.)
|
|
// Not stored in Session (not session-cumulative); used only for the final
|
|
// end-of-turn wire emission.
|
|
let mut turn_pricing_identity: Option<Option<crate::types::PricingIdentity>> = None;
|
|
let mut ctx = RunCtx {
|
|
cfg: &app.cfg,
|
|
effective_model: effective_model_str,
|
|
session_id: &sid,
|
|
system_prompt: &effective_system_prompt,
|
|
llm: &app.llm,
|
|
mcp: &mcp,
|
|
skills: &skills,
|
|
wire: &wire_tx,
|
|
cancel: &mut cancel_rx,
|
|
steer: &mut steer_rx,
|
|
history: &mut history,
|
|
original_task: &mut original_task,
|
|
handoff_count: &mut handoff_count,
|
|
run_id,
|
|
last_request_input_tokens: &mut last_request_input_tokens,
|
|
last_request_history_bytes: &mut last_request_history_bytes,
|
|
turn_input_tokens: &mut turn_input_tokens,
|
|
turn_output_tokens: &mut turn_output_tokens,
|
|
turn_cached_input_tokens: &mut turn_cached_input_tokens,
|
|
turn_cache_write_tokens: &mut turn_cache_write_tokens,
|
|
turn_total_state: &mut turn_total_state,
|
|
turn_pricing_identity: &mut turn_pricing_identity,
|
|
usage_baseline,
|
|
};
|
|
let result = ctx.run(p.prompt).await;
|
|
if let Some(s) = app.sessions.lock().await.get_mut(&sid) {
|
|
s.busy = false;
|
|
// Clear run state so a late steer can't queue into a finished turn.
|
|
s.active_run_id = None;
|
|
s.steer_tx = None;
|
|
s.history = history;
|
|
s.original_task = original_task;
|
|
s.handoff_count = handoff_count;
|
|
s.last_request_input_tokens = last_request_input_tokens;
|
|
s.last_request_history_bytes = last_request_history_bytes;
|
|
}
|
|
// Update session-cumulative token counters and emit the usage notification
|
|
// BEFORE sending the session/prompt response. buzz-acp's UsageTracker
|
|
// processes the notification while the turn is still in-flight (i.e. before
|
|
// the response triggers take_turn_usage()), which is required for the
|
|
// begin_turn gate to recognise it as publishable.
|
|
//
|
|
// Only emit when at least one token count was observed — a turn with no
|
|
// provider response (validation failure, pre-response cancellation) carries
|
|
// no information and must not produce a kind 44200 record per NIP-AM.
|
|
if !matches!(turn_input_tokens, crate::types::TurnIOState::Unseen)
|
|
|| !matches!(turn_output_tokens, crate::types::TurnIOState::Unseen)
|
|
{
|
|
let accumulated = {
|
|
let mut sessions = app.sessions.lock().await;
|
|
if let Some(s) = sessions.get_mut(&sid) {
|
|
// merge_session: Poisoned poisons permanently; Exact sums with
|
|
// overflow-check → Poisoned on wrap; Unseen leaves unchanged.
|
|
s.accumulated_input_tokens =
|
|
s.accumulated_input_tokens.merge_session(turn_input_tokens);
|
|
s.accumulated_output_tokens = s
|
|
.accumulated_output_tokens
|
|
.merge_session(turn_output_tokens);
|
|
// D1 tri-state merge: merge_session propagates Unknown when
|
|
// the turn was poisoned (any usage-bearing round omitted the
|
|
// category), and is a no-op when the turn was Unseen (no
|
|
// usage-bearing response at all).
|
|
s.accumulated_cached_input_tokens = s
|
|
.accumulated_cached_input_tokens
|
|
.merge_session(turn_cached_input_tokens);
|
|
s.accumulated_cache_write_tokens = s
|
|
.accumulated_cache_write_tokens
|
|
.merge_session(turn_cache_write_tokens);
|
|
// Fold the per-turn total state into the session cumulative.
|
|
// Unknown poisons the session permanently; Exact adds to running sum;
|
|
// Unseen (turn emitted no usage) leaves the cumulative unchanged.
|
|
// Uses TurnTotalState::merge_session, which applies the same
|
|
// checked-add / overflow-poisons contract as the per-response fold.
|
|
s.accumulated_total_state =
|
|
s.accumulated_total_state.merge_session(turn_total_state);
|
|
Some((
|
|
s.accumulated_input_tokens,
|
|
s.accumulated_output_tokens,
|
|
s.accumulated_cached_input_tokens,
|
|
s.accumulated_cache_write_tokens,
|
|
s.accumulated_total_state,
|
|
))
|
|
} else {
|
|
// Session is gone — the accumulated baseline no longer exists, so
|
|
// there is nothing correct to emit. Skip the usage notification.
|
|
None
|
|
}
|
|
};
|
|
if let Some((
|
|
accumulated_in,
|
|
accumulated_out,
|
|
accumulated_cached,
|
|
accumulated_written,
|
|
accumulated_total,
|
|
)) = accumulated
|
|
{
|
|
// Same builder the run loop uses for its per-round reports, so the
|
|
// final notification is shape-identical to the ones that preceded
|
|
// it and a consumer taking the high-water mark lands on this one.
|
|
let update = wire::usage_update_payload(
|
|
accumulated_in.exact_value(),
|
|
accumulated_out.exact_value(),
|
|
accumulated_cached.exact_value(),
|
|
accumulated_written.exact_value(),
|
|
accumulated_total,
|
|
effective_model_str,
|
|
// Pass the proven per-turn identity if consistent; absent otherwise.
|
|
turn_pricing_identity
|
|
.as_ref()
|
|
.and_then(|inner| inner.as_ref()),
|
|
);
|
|
wire::send(&wire_tx, goose_session_update(&sid, update)).await;
|
|
}
|
|
}
|
|
match result {
|
|
Ok(stop) => {
|
|
wire::send(
|
|
&wire_tx,
|
|
wire::ok(id, json!({ "stopReason": stop.as_wire() })),
|
|
)
|
|
.await
|
|
}
|
|
Err(e) => wire::send(&wire_tx, wire::err(id, e.json_rpc_code(), &e.to_string())).await,
|
|
}
|
|
}
|
|
|
|
async fn acquire_session(
|
|
app: &Arc<App>,
|
|
session_id: &str,
|
|
) -> Result<
|
|
(
|
|
String,
|
|
Arc<McpRegistry>,
|
|
Vec<SkillEntry>,
|
|
Vec<HistoryItem>,
|
|
Option<String>,
|
|
usize,
|
|
Option<u64>,
|
|
Option<usize>,
|
|
watch::Receiver<bool>,
|
|
Arc<str>,
|
|
Option<String>,
|
|
String,
|
|
mpsc::UnboundedReceiver<Vec<ContentBlock>>,
|
|
crate::types::SessionUsageBaseline,
|
|
),
|
|
&'static str,
|
|
> {
|
|
let mut sessions = app.sessions.lock().await;
|
|
let s = sessions.get_mut(session_id).ok_or("unknown session")?;
|
|
if s.busy {
|
|
return Err("prompt already in flight");
|
|
}
|
|
// Generate the run id before mutating session state. On RNG failure we reject
|
|
// the prompt cleanly: the session stays idle and the caller can retry. Generating
|
|
// after `s.busy = true` with `?` would wedge the session permanently busy.
|
|
let run_id = format!(
|
|
"run_{}",
|
|
session_token().map_err(|_| "rng failure; retry prompt")?
|
|
);
|
|
s.busy = true;
|
|
let (tx, rx) = watch::channel(false);
|
|
s.cancel_tx = tx;
|
|
// Skills are read-only after session creation; clone the Vec so RunCtx
|
|
// can hold a reference without holding the sessions lock.
|
|
let skills = s.skills.clone();
|
|
// Fresh run id + steer channel for this turn. The run id lets steer-capable
|
|
// clients target *this* turn (rejecting steers aimed at a turn that already
|
|
// ended); the channel carries mid-turn injections to the run loop.
|
|
s.active_run_id = Some(run_id.clone());
|
|
let (steer_tx, steer_rx) = mpsc::unbounded_channel();
|
|
s.steer_tx = Some(steer_tx);
|
|
let effective_model = s.effective_model.clone();
|
|
Ok((
|
|
s.id.clone(),
|
|
s.mcp.clone(),
|
|
skills,
|
|
std::mem::take(&mut s.history),
|
|
s.original_task.take(),
|
|
s.handoff_count,
|
|
s.last_request_input_tokens,
|
|
s.last_request_history_bytes,
|
|
rx,
|
|
Arc::clone(&s.effective_system_prompt),
|
|
effective_model,
|
|
run_id,
|
|
steer_rx,
|
|
// Snapshot rather than a handle: the run loop reports cumulative usage
|
|
// after every LLM round, and taking the sessions lock on each of those
|
|
// would serialise concurrent sessions behind one another's provider
|
|
// round-trips. Nothing else advances these counters while this turn
|
|
// holds `busy`, so the snapshot cannot go stale under it.
|
|
crate::types::SessionUsageBaseline {
|
|
input_tokens: s.accumulated_input_tokens,
|
|
output_tokens: s.accumulated_output_tokens,
|
|
cached_input_tokens: s.accumulated_cached_input_tokens,
|
|
cache_write_tokens: s.accumulated_cache_write_tokens,
|
|
total_state: s.accumulated_total_state,
|
|
},
|
|
))
|
|
}
|
|
|
|
fn session_token() -> Result<String, String> {
|
|
let mut b = [0u8; 8];
|
|
getrandom::fill(&mut b).map_err(|e| format!("rng: getrandom failed: {e}"))?;
|
|
Ok(b.iter().map(|x| format!("{x:02x}")).collect())
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use crate::catalog::ModelEntry;
|
|
use crate::types::AgentError;
|
|
|
|
/// Regression: a discovery error must not pin the models_cache for the process lifetime.
|
|
///
|
|
/// `resolve_models_catalog` uses `get_or_try_init` so an `Err` leaves the `OnceCell`
|
|
/// empty and the next `session/new` retries discovery. This test calls
|
|
/// `resolve_models_catalog` directly — the same function `session_new` calls — so
|
|
/// reverting `session_new` to `get_or_init` (or any other cache-on-error variant) would
|
|
/// break this test, not just the standalone `OnceCell` semantics.
|
|
#[tokio::test]
|
|
async fn models_cache_does_not_pin_on_discovery_error() {
|
|
let cache: tokio::sync::OnceCell<Vec<ModelEntry>> = tokio::sync::OnceCell::new();
|
|
|
|
// First call — discovery failure is surfaced and leaves the cell empty.
|
|
let error = crate::resolve_models_catalog(&cache, async {
|
|
Err::<Vec<ModelEntry>, AgentError>(AgentError::Llm("transient failure".into()))
|
|
})
|
|
.await
|
|
.unwrap_err();
|
|
assert!(matches!(error, AgentError::Llm(_)));
|
|
|
|
// Second call — discovery succeeds. Cell is now populated and returned.
|
|
let discovered = vec![ModelEntry {
|
|
id: "databricks-meta-llama-3-1-70b-instruct".into(),
|
|
name: "databricks-meta-llama-3-1-70b-instruct".into(),
|
|
}];
|
|
let discovered_clone = discovered.clone();
|
|
let second = crate::resolve_models_catalog(&cache, async move {
|
|
Ok::<Vec<ModelEntry>, AgentError>(discovered_clone)
|
|
})
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(
|
|
second, discovered,
|
|
"second call must return the discovered catalog"
|
|
);
|
|
assert!(
|
|
cache.get().is_some(),
|
|
"cell must be populated after successful discovery"
|
|
);
|
|
assert_eq!(
|
|
cache.get().unwrap(),
|
|
&discovered,
|
|
"cache must hold the successful discovery result"
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn models_catalog_does_not_cache_oauth_auth_fallback() {
|
|
let cache: tokio::sync::OnceCell<Vec<ModelEntry>> = tokio::sync::OnceCell::new();
|
|
let error = crate::resolve_models_catalog(&cache, async {
|
|
Err::<Vec<ModelEntry>, AgentError>(AgentError::LlmAuth("sign in again".into()))
|
|
})
|
|
.await
|
|
.unwrap_err();
|
|
|
|
assert!(matches!(error, AgentError::LlmAuth(_)));
|
|
assert!(cache.get().is_none());
|
|
|
|
let discovered = vec![ModelEntry {
|
|
id: "authenticated-model".into(),
|
|
name: "authenticated-model".into(),
|
|
}];
|
|
let result = crate::resolve_models_catalog(&cache, async {
|
|
Ok::<Vec<ModelEntry>, AgentError>(discovered.clone())
|
|
})
|
|
.await
|
|
.unwrap();
|
|
|
|
assert_eq!(result, discovered);
|
|
assert_eq!(cache.get(), Some(&discovered));
|
|
}
|
|
|
|
#[test]
|
|
fn configured_model_fallback_is_trimmed_and_singular() {
|
|
assert_eq!(
|
|
crate::configured_model_fallback(" configured-model "),
|
|
vec![ModelEntry {
|
|
id: "configured-model".into(),
|
|
name: "configured-model".into(),
|
|
}]
|
|
);
|
|
}
|
|
}
|