mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-10-10 03:30:59 +00:00
Merge pull request #598 from fabro-sh/run-session-trace-header
feat(llm): send x-session-id trace header with the run ID
This commit is contained in:
commit
21e84484d2
6 changed files with 220 additions and 16 deletions
175
lib/crates/fabro-auth/src/extra_headers_source.rs
Normal file
175
lib/crates/fabro-auth/src/extra_headers_source.rs
Normal file
|
|
@ -0,0 +1,175 @@
|
|||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use fabro_model::{Catalog, ProviderId};
|
||||
|
||||
use crate::credential_source::{CredentialSource, ResolvedCredentials};
|
||||
|
||||
/// Decorates another [`CredentialSource`] by appending fixed extra headers to
|
||||
/// every credential it resolves.
|
||||
///
|
||||
/// Headers already present on a credential (for example from explicit
|
||||
/// provider configuration) are left untouched.
|
||||
pub struct ExtraHeadersCredentialSource {
|
||||
inner: Arc<dyn CredentialSource>,
|
||||
headers: HashMap<String, String>,
|
||||
}
|
||||
|
||||
impl ExtraHeadersCredentialSource {
|
||||
#[must_use]
|
||||
pub fn new(inner: Arc<dyn CredentialSource>, headers: HashMap<String, String>) -> Self {
|
||||
Self { inner, headers }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl CredentialSource for ExtraHeadersCredentialSource {
|
||||
async fn resolve(&self, catalog: &Catalog) -> anyhow::Result<ResolvedCredentials> {
|
||||
let mut resolved = self.inner.resolve(catalog).await?;
|
||||
for credential in &mut resolved.credentials {
|
||||
for (name, value) in &self.headers {
|
||||
if credential
|
||||
.extra_headers
|
||||
.keys()
|
||||
.any(|existing| existing.eq_ignore_ascii_case(name))
|
||||
{
|
||||
continue;
|
||||
}
|
||||
credential.extra_headers.insert(name.clone(), value.clone());
|
||||
}
|
||||
}
|
||||
Ok(resolved)
|
||||
}
|
||||
|
||||
async fn configured_providers(&self, catalog: &Catalog) -> Vec<ProviderId> {
|
||||
self.inner.configured_providers(catalog).await
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::{ApiCredential, ResolveError};
|
||||
|
||||
struct StubSource {
|
||||
credentials: Vec<ApiCredential>,
|
||||
auth_issue_provider: Option<ProviderId>,
|
||||
configured_providers: Vec<ProviderId>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl CredentialSource for StubSource {
|
||||
async fn resolve(&self, _catalog: &Catalog) -> anyhow::Result<ResolvedCredentials> {
|
||||
Ok(ResolvedCredentials {
|
||||
credentials: self.credentials.clone(),
|
||||
auth_issues: self
|
||||
.auth_issue_provider
|
||||
.iter()
|
||||
.map(|provider| {
|
||||
(
|
||||
provider.clone(),
|
||||
ResolveError::RefreshTokenMissing(provider.clone()),
|
||||
)
|
||||
})
|
||||
.collect(),
|
||||
})
|
||||
}
|
||||
|
||||
async fn configured_providers(&self, _catalog: &Catalog) -> Vec<ProviderId> {
|
||||
self.configured_providers.clone()
|
||||
}
|
||||
}
|
||||
|
||||
fn credential(provider: ProviderId, extra_headers: HashMap<String, String>) -> ApiCredential {
|
||||
ApiCredential {
|
||||
provider,
|
||||
auth_header: None,
|
||||
extra_headers,
|
||||
base_url: None,
|
||||
codex_mode: false,
|
||||
org_id: None,
|
||||
project_id: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn appends_headers_to_every_resolved_credential() {
|
||||
let source = ExtraHeadersCredentialSource::new(
|
||||
Arc::new(StubSource {
|
||||
credentials: vec![
|
||||
credential(ProviderId::anthropic(), HashMap::new()),
|
||||
credential(ProviderId::openai(), HashMap::new()),
|
||||
],
|
||||
auth_issue_provider: None,
|
||||
configured_providers: Vec::new(),
|
||||
}),
|
||||
HashMap::from([("x-session-id".to_string(), "run-123".to_string())]),
|
||||
);
|
||||
|
||||
let resolved = source.resolve(Catalog::builtin()).await.unwrap();
|
||||
|
||||
assert_eq!(resolved.credentials.len(), 2);
|
||||
for credential in &resolved.credentials {
|
||||
assert_eq!(
|
||||
credential
|
||||
.extra_headers
|
||||
.get("x-session-id")
|
||||
.map(String::as_str),
|
||||
Some("run-123")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn preserves_case_insensitive_headers_already_set_on_a_credential() {
|
||||
let source = ExtraHeadersCredentialSource::new(
|
||||
Arc::new(StubSource {
|
||||
credentials: vec![credential(
|
||||
ProviderId::new("openrouter"),
|
||||
HashMap::from([("X-Session-Id".to_string(), "configured".to_string())]),
|
||||
)],
|
||||
auth_issue_provider: None,
|
||||
configured_providers: Vec::new(),
|
||||
}),
|
||||
HashMap::from([("x-session-id".to_string(), "run-123".to_string())]),
|
||||
);
|
||||
|
||||
let resolved = source.resolve(Catalog::builtin()).await.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
resolved.credentials[0]
|
||||
.extra_headers
|
||||
.get("X-Session-Id")
|
||||
.map(String::as_str),
|
||||
Some("configured")
|
||||
);
|
||||
assert_eq!(resolved.credentials[0].extra_headers.len(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn passes_through_auth_issues_and_configured_providers() {
|
||||
let auth_issue_provider = ProviderId::anthropic();
|
||||
let configured_provider = ProviderId::gemini();
|
||||
let source = ExtraHeadersCredentialSource::new(
|
||||
Arc::new(StubSource {
|
||||
credentials: vec![credential(ProviderId::openai(), HashMap::new())],
|
||||
auth_issue_provider: Some(auth_issue_provider.clone()),
|
||||
configured_providers: vec![configured_provider.clone()],
|
||||
}),
|
||||
HashMap::from([("x-session-id".to_string(), "run-123".to_string())]),
|
||||
);
|
||||
|
||||
let resolved = source.resolve(Catalog::builtin()).await.unwrap();
|
||||
let [(reported_provider, ResolveError::RefreshTokenMissing(error_provider))] =
|
||||
resolved.auth_issues.as_slice()
|
||||
else {
|
||||
panic!("expected the inner source's refresh-token issue");
|
||||
};
|
||||
assert_eq!(reported_provider, &auth_issue_provider);
|
||||
assert_eq!(error_provider, &auth_issue_provider);
|
||||
|
||||
let providers = source.configured_providers(Catalog::builtin()).await;
|
||||
assert_eq!(providers, vec![configured_provider]);
|
||||
}
|
||||
}
|
||||
|
|
@ -2,6 +2,7 @@ mod context;
|
|||
mod credential;
|
||||
mod credential_source;
|
||||
mod env_source;
|
||||
mod extra_headers_source;
|
||||
mod refresh;
|
||||
mod resolve;
|
||||
mod sql_vault_source;
|
||||
|
|
@ -15,6 +16,7 @@ pub use context::{AuthContextRequest, AuthContextResponse};
|
|||
pub use credential::{ApiKeyHeader, OAuthConfig, OAuthCredential, OAuthTokens};
|
||||
pub use credential_source::{CredentialSource, ResolvedCredentials};
|
||||
pub use env_source::EnvCredentialSource;
|
||||
pub use extra_headers_source::ExtraHeadersCredentialSource;
|
||||
pub use refresh::refresh_oauth_credential;
|
||||
pub use resolve::{
|
||||
ApiCredential, CredentialResolver, CredentialUsage, EnvLookup, ResolveError,
|
||||
|
|
|
|||
|
|
@ -858,7 +858,6 @@ impl RunSession {
|
|||
store_progress_logger.register(self.emitter.as_ref());
|
||||
|
||||
let init_options = InitOptions {
|
||||
run_id: record.run_id,
|
||||
run_store: self.run_store.clone(),
|
||||
dry_run: run_options.dry_run_enabled(),
|
||||
emitter: self.emitter,
|
||||
|
|
|
|||
|
|
@ -256,7 +256,6 @@ async fn execute_test_run_with_options(
|
|||
let initialized = initialize(
|
||||
persisted_workflow(graph, String::new(), &run_options.run_dir, run_id_value),
|
||||
InitOptions {
|
||||
run_id: run_id_value,
|
||||
run_store: run_store.into(),
|
||||
dry_run: false,
|
||||
emitter: emitter.clone(),
|
||||
|
|
@ -317,7 +316,6 @@ async fn execute_runs_start_to_exit_and_returns_final_context() {
|
|||
let initialized = initialize(
|
||||
persisted_workflow(graph, source, &run_dir, test_run_id("run-test")),
|
||||
InitOptions {
|
||||
run_id: test_run_id("run-test"),
|
||||
run_store: run_store.into(),
|
||||
dry_run: false,
|
||||
emitter: test_emitter_arc("run-test"),
|
||||
|
|
@ -393,7 +391,6 @@ async fn run_with_lifecycle(
|
|||
let initialized = initialize(
|
||||
persisted_workflow(graph.clone(), String::new(), &run_dir, run_id),
|
||||
InitOptions {
|
||||
run_id,
|
||||
run_store: run_store.into(),
|
||||
dry_run: false,
|
||||
emitter: emitter.clone(),
|
||||
|
|
|
|||
|
|
@ -5,7 +5,8 @@ use std::time::Instant;
|
|||
|
||||
use fabro_agent::{Sandbox, ToolSecrets};
|
||||
use fabro_auth::{
|
||||
CredentialSource, EnvCredentialSource, VaultCredentialSource, auth_issue_message,
|
||||
CredentialSource, EnvCredentialSource, ExtraHeadersCredentialSource, VaultCredentialSource,
|
||||
auth_issue_message,
|
||||
};
|
||||
use fabro_graphviz::graph;
|
||||
use fabro_hooks::{HookContext, HookDecision, HookEvent, HookExecutionContext, HookRunner};
|
||||
|
|
@ -260,11 +261,23 @@ fn graph_needs_api_backend(graph: &graph::Graph) -> bool {
|
|||
graph.nodes.values().any(routing::node_needs_api_backend)
|
||||
}
|
||||
|
||||
fn build_llm_source(vault: Option<Arc<AsyncRwLock<Vault>>>) -> Arc<dyn CredentialSource> {
|
||||
match vault {
|
||||
/// Trace header attached to every LLM request in a run so gateways that
|
||||
/// understand it (e.g. OpenRouter broadcast) can group the run's requests
|
||||
/// into one session. Explicit `extra_headers` provider configuration wins.
|
||||
const SESSION_ID_HEADER: &str = "x-session-id";
|
||||
|
||||
fn build_llm_source(
|
||||
vault: Option<Arc<AsyncRwLock<Vault>>>,
|
||||
run_id: fabro_types::RunId,
|
||||
) -> Arc<dyn CredentialSource> {
|
||||
let inner: Arc<dyn CredentialSource> = match vault {
|
||||
Some(vault) => Arc::new(VaultCredentialSource::new(vault)),
|
||||
None => Arc::new(EnvCredentialSource::new()),
|
||||
}
|
||||
};
|
||||
Arc::new(ExtraHeadersCredentialSource::new(
|
||||
inner,
|
||||
HashMap::from([(SESSION_ID_HEADER.to_string(), run_id.to_string())]),
|
||||
))
|
||||
}
|
||||
|
||||
/// INITIALIZE phase: prepare the sandbox, env, and handlers for execution.
|
||||
|
|
@ -277,7 +290,7 @@ pub async fn initialize(
|
|||
options.run_options.run_dir = run_dir.clone();
|
||||
options.run_options.git = options.git.clone();
|
||||
|
||||
let llm_source = build_llm_source(options.vault.clone());
|
||||
let llm_source = build_llm_source(options.vault.clone(), options.run_options.run_id);
|
||||
let tool_secrets = tool_secrets_from_configured_sources(options.vault.as_ref()).await;
|
||||
let catalog = Arc::clone(&options.catalog);
|
||||
let sandbox_git = Arc::new(SandboxGitRuntime::new());
|
||||
|
|
@ -347,7 +360,7 @@ pub async fn initialize(
|
|||
let sandbox = reconnect_for_run_with_callback(
|
||||
instance,
|
||||
daytona_api_key,
|
||||
Some(options.run_id),
|
||||
Some(options.run_options.run_id),
|
||||
Some(Arc::clone(&sandbox_event_callback)),
|
||||
)
|
||||
.await
|
||||
|
|
@ -812,7 +825,6 @@ mod tests {
|
|||
});
|
||||
|
||||
let result = initialize(persisted, InitOptions {
|
||||
run_id: test_run_id(),
|
||||
run_store: {
|
||||
let store = memory_store();
|
||||
let inner = store.create_run(&test_run_id()).await.unwrap();
|
||||
|
|
@ -894,7 +906,6 @@ mod tests {
|
|||
let emitter = Arc::new(crate::event::Emitter::new(test_run_id()));
|
||||
|
||||
let initialized = initialize(persisted, InitOptions {
|
||||
run_id: test_run_id(),
|
||||
run_store: {
|
||||
let store = memory_store();
|
||||
let inner = store.create_run(&test_run_id()).await.unwrap();
|
||||
|
|
@ -1028,6 +1039,30 @@ mod tests {
|
|||
assert!(!effective_dry_run);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn build_llm_source_appends_run_session_trace_header() {
|
||||
let mut vault = Vault::from_entries(HashMap::new());
|
||||
fabro_auth::vault_set_token(&mut vault, EnvVars::ANTHROPIC_API_KEY, "anthropic-key")
|
||||
.unwrap();
|
||||
let vault = Arc::new(AsyncRwLock::new(vault));
|
||||
let run_id = test_run_id();
|
||||
let expected_session_id = run_id.to_string();
|
||||
|
||||
let source = build_llm_source(Some(vault), run_id);
|
||||
let resolved = source.resolve(test_catalog().as_ref()).await.unwrap();
|
||||
|
||||
assert!(!resolved.credentials.is_empty());
|
||||
for credential in &resolved.credentials {
|
||||
assert_eq!(
|
||||
credential
|
||||
.extra_headers
|
||||
.get(SESSION_ID_HEADER)
|
||||
.map(String::as_str),
|
||||
Some(expected_session_id.as_str())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn initialize_executes_acp_backend_node_from_registry() {
|
||||
let temp = tempfile::tempdir().unwrap();
|
||||
|
|
@ -1096,7 +1131,6 @@ mod tests {
|
|||
let store = memory_store();
|
||||
let run_store = store.create_run(&test_run_id()).await.unwrap();
|
||||
let initialized = initialize(test_persisted(graph, source, &run_dir), InitOptions {
|
||||
run_id: test_run_id(),
|
||||
run_store: run_store.into(),
|
||||
dry_run: false,
|
||||
emitter: emitter.clone(),
|
||||
|
|
@ -1192,7 +1226,6 @@ mod tests {
|
|||
store_logger.register(&emitter);
|
||||
|
||||
let initialized = initialize(persisted, InitOptions {
|
||||
run_id: test_run_id(),
|
||||
run_store: run_store.into(),
|
||||
dry_run: false,
|
||||
emitter: emitter.clone(),
|
||||
|
|
@ -1331,7 +1364,6 @@ mod tests {
|
|||
|
||||
let emitter = Arc::new(crate::event::Emitter::new(test_run_id()));
|
||||
let result = initialize(persisted, InitOptions {
|
||||
run_id: test_run_id(),
|
||||
run_store: {
|
||||
let store = memory_store();
|
||||
let inner = store.create_run(&test_run_id()).await.unwrap();
|
||||
|
|
|
|||
|
|
@ -249,7 +249,6 @@ pub struct SandboxEnvSpec {
|
|||
}
|
||||
|
||||
pub struct InitOptions {
|
||||
pub run_id: RunId,
|
||||
pub run_store: RunStoreHandle,
|
||||
pub dry_run: bool,
|
||||
pub emitter: Arc<Emitter>,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue