fix(memory): scope embedding credentials to the endpoint that may receive them

`MemoryBackendParams` carried the primary session `AuthManager` and the
session api-key provider unconditionally, while `embed_base_url` is the
CURRENT MODEL's endpoint and `AuthRetryMiddleware` stamps `Authorization`
on every request it wraps. A user who enables `[memory.embedding] model`
while running a BYOK or subscription-OAuth model therefore sent the Kimi
session bearer to that third party.

`EndpointScopedCredentials` binds the credential to the one endpoint it
may reach: `for_endpoint` drops the handle unless the caller vouches for
the URL, and `approved_for` re-checks at provider-build time in release
too, because `MemoryBackendParams` is `Clone` and callers rewrite fields
on the copy.

The shell decides through `CredentialAuthority::manager_for` rather than
a second URL predicate — it answers both whether a credential may ride
and which manager governs it, so a subscription-OAuth platform gets its
own pooled manager. The session's `SharedApiKeyProvider` is not
forwarded at all: it is hard-wired to the primary manager, so at a
pooled platform's host it would resolve the wrong bearer. A platform's
own `embed_api_key` is untouched and keeps serving its own endpoint.

The background reindex built a second provider straight from
`ApiEmbeddingProvider::from_session`, outside the chokepoint and without
401 refresh; it now embeds through the session's own params.

Test strength verified by mutation: with the guard reverted, exactly
`session_credentials_are_withheld_from_a_foreign_endpoint` and
`a_cloned_param_set_cannot_redirect_scoped_credentials` fail.
This commit is contained in:
2026-07-26 23:41:45 -04:00
parent 867b3e110b
commit ed8049cf77
7 changed files with 295 additions and 109 deletions
+203 -77
View File
@@ -18,6 +18,76 @@ use super::embedding::EmbeddingProvider as _;
use super::storage::MemoryStorage;
use super::watcher::MemoryFileWatcher;
/// The session's embedding credentials, bound to the one endpoint they may
/// reach. Only [`Self::for_endpoint`] retains a live handle; the default fails
/// closed.
///
/// `embed_base_url` is the CURRENT MODEL's endpoint, so a session on a BYOK or
/// subscription-OAuth model aims memory embeddings at that provider's host —
/// and [`kigi_auth::AuthRetryMiddleware`] stamps `Authorization` on every
/// request it wraps, with no idea where the request is going. The caller that
/// owns the credential rule decides `trusted` once, here; a platform's own
/// `embed_api_key` is unaffected and keeps serving its own endpoint.
#[derive(Clone, Default)]
pub struct EndpointScopedCredentials {
endpoint: Option<reqwest::Url>,
auth_credentials: Option<Arc<dyn kigi_auth::AuthCredentialProvider>>,
}
// Redacts the credential handles; only their presence is printable.
impl std::fmt::Debug for EndpointScopedCredentials {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("EndpointScopedCredentials")
.field("endpoint", &self.endpoint)
.field("has_auth_credentials", &self.auth_credentials.is_some())
.finish()
}
}
impl EndpointScopedCredentials {
/// No session credential may ride — the state every non-session caller wants.
pub fn none() -> Self {
Self::default()
}
/// Retains the handles only for a `trusted`, parsable `endpoint`.
pub fn for_endpoint(
endpoint: &str,
trusted: bool,
auth_credentials: Option<Arc<dyn kigi_auth::AuthCredentialProvider>>,
) -> Self {
if trusted && let Ok(url) = reqwest::Url::parse(endpoint) {
return Self {
endpoint: Some(url),
auth_credentials,
};
}
if auth_credentials.is_some() {
tracing::info!(
target: kigi_log::memory_log::TARGET,
endpoint,
"memory embeddings: session credentials withheld from this endpoint; \
its own key, if any, still applies"
);
}
Self::none()
}
pub fn is_empty(&self) -> bool {
self.auth_credentials.is_none()
}
/// Enforced at request-build time, in release too: [`MemoryBackendParams`]
/// is `Clone` and callers rewrite fields on the copy, so construction-time
/// scoping alone would not survive a rewritten `embed_base_url`.
fn approved_for(&self, base_url: &str) -> bool {
match &self.endpoint {
None => self.is_empty(),
Some(endpoint) => reqwest::Url::parse(base_url).is_ok_and(|url| &url == endpoint),
}
}
}
/// All configuration needed to build a fully-wired [`MemoryBackendImpl`] for a live session.
///
/// Grouping these in one struct ensures every call site — ToolBridge, first-turn
@@ -30,7 +100,8 @@ pub struct MemoryBackendParams {
pub session_id: String,
/// Embedding provider config — `None` forces FTS-only fallback everywhere.
pub embed_config: Option<kigi_config_types::MemoryEmbeddingConfig>,
/// Base URL for embedding API calls (CLI proxy).
/// Base URL for embedding API calls (CLI proxy). Must match the endpoint
/// `embedding_credentials` was scoped to; a mismatch fails closed.
pub embed_base_url: String,
/// API key for embedding API calls.
pub embed_api_key: Option<String>,
@@ -47,10 +118,8 @@ pub struct MemoryBackendParams {
/// - `"injection"` — first-turn memory context injection
/// - `"compaction_recovery"` — post-compaction context re-injection
pub search_source: &'static str,
/// Dynamic API key provider — when set, `make_embedding_provider()` resolves
/// the key per-call instead of using the static `embed_api_key`.
pub api_key_provider: Option<kigi_tools::types::SharedApiKeyProvider>,
pub auth_credentials: Option<Arc<dyn kigi_auth::AuthCredentialProvider>>,
/// The session credentials, and the single endpoint they may reach.
pub embedding_credentials: EndpointScopedCredentials,
}
impl MemoryBackendParams {
@@ -59,8 +128,7 @@ impl MemoryBackendParams {
pub async fn make_embedding_provider(&self) -> Option<super::embedding::ApiEmbeddingProvider> {
build_embedding_provider(
self.embed_config.as_ref(),
self.auth_credentials.as_ref(),
self.api_key_provider.as_ref(),
&self.embedding_credentials,
self.embed_api_key.as_deref(),
&self.embed_base_url,
)
@@ -70,8 +138,7 @@ impl MemoryBackendParams {
async fn build_embedding_provider(
config: Option<&kigi_config_types::MemoryEmbeddingConfig>,
auth_credentials: Option<&Arc<dyn kigi_auth::AuthCredentialProvider>>,
api_key_provider: Option<&kigi_tools::types::SharedApiKeyProvider>,
credentials: &EndpointScopedCredentials,
static_api_key: Option<&str>,
base_url: &str,
) -> Option<super::embedding::ApiEmbeddingProvider> {
@@ -80,9 +147,19 @@ async fn build_embedding_provider(
return None;
}
let approved = credentials.approved_for(base_url);
if !approved {
tracing::error!(
target: kigi_log::memory_log::TARGET,
base_url,
approved_endpoint = ?credentials.endpoint,
"memory embeddings: scoped credentials do not match the request URL; dropping them"
);
}
// Prefer the refresh-capable credential provider — the middleware gives
// 401 retry for free without any per-call key resolution.
if let Some(creds) = auth_credentials {
if approved && let Some(creds) = credentials.auth_credentials.as_ref() {
let client = super::embedding::build_middleware_client(creds.clone());
return super::embedding::ApiEmbeddingProvider::from_config(
config,
@@ -91,14 +168,13 @@ async fn build_embedding_provider(
);
}
// Fallback: resolve API key per-call, wrap in a static middleware client
// (no 401 refresh, but auth header is still stamped by middleware).
let api_key = match api_key_provider {
Some(p) => p.current_api_key_async().await,
None => None,
}
.or_else(|| static_api_key.map(|s| s.to_owned()))?;
super::embedding::ApiEmbeddingProvider::from_session(config, base_url.to_owned(), api_key)
// The platform's own configured key, wrapped in a static middleware client
// (no 401 refresh, but the auth header is still stamped by middleware).
super::embedding::ApiEmbeddingProvider::from_session(
config,
base_url.to_owned(),
static_api_key?.to_owned(),
)
}
/// `MemoryBackend` implementation backed by hybrid search (FTS5 + vector KNN).
@@ -129,10 +205,8 @@ pub struct MemoryBackendImpl {
/// Only the ToolBridge backend's counter is shared back to the session actor;
/// injection and compaction-recovery backends use their own local counters.
pub search_counter: std::sync::Arc<std::sync::atomic::AtomicU64>,
/// Dynamic API key provider for embedding requests.
api_key_provider: Option<kigi_tools::types::SharedApiKeyProvider>,
/// Refresh-capable credential provider for embedding HTTP middleware.
auth_credentials: Option<Arc<dyn kigi_auth::AuthCredentialProvider>>,
/// The session credentials, and the single endpoint they may reach.
embedding_credentials: EndpointScopedCredentials,
}
impl MemoryBackendImpl {
@@ -150,8 +224,7 @@ impl MemoryBackendImpl {
stale_claim_secs: 60,
session_id: String::new(),
search_source: "tool",
api_key_provider: None,
auth_credentials: None,
embedding_credentials: EndpointScopedCredentials::none(),
search_counter: std::sync::Arc::new(std::sync::atomic::AtomicU64::new(0)),
}
}
@@ -200,8 +273,7 @@ impl MemoryBackendImpl {
async fn make_embedding_provider(&self) -> Option<super::embedding::ApiEmbeddingProvider> {
build_embedding_provider(
self.embed_config.as_ref(),
self.auth_credentials.as_ref(),
self.api_key_provider.as_ref(),
&self.embedding_credentials,
self.embed_api_key.as_deref(),
&self.embed_base_url,
)
@@ -232,8 +304,7 @@ impl MemoryBackendImpl {
if let Some(w) = &params.watcher {
backend = backend.with_watcher(w.clone(), params.stale_claim_secs);
}
backend.api_key_provider = params.api_key_provider.clone();
backend.auth_credentials = params.auth_credentials.clone();
backend.embedding_credentials = params.embedding_credentials.clone();
backend
}
}
@@ -479,8 +550,7 @@ mod factory_tests {
watcher: None,
stale_claim_secs: 60,
search_source: "tool",
api_key_provider: None,
auth_credentials: None,
embedding_credentials: EndpointScopedCredentials::none(),
}
}
@@ -1038,72 +1108,128 @@ mod factory_tests {
);
}
/// Regression: provider build must use `current_api_key_async`,
/// never sync. Prevents memory_search 401s on rotated tokens.
#[tokio::test]
async fn make_embedding_provider_uses_async_api_key_resolution() {
use kigi_tools::types::ApiKeyProvider;
use std::sync::atomic::{AtomicU32, Ordering};
struct AsyncProbe {
sync_calls: Arc<AtomicU32>,
async_calls: Arc<AtomicU32>,
struct ProbeCredentials;
impl kigi_auth::HttpAuth for ProbeCredentials {
fn apply(
&self,
builder: reqwest::RequestBuilder,
_base_url: &str,
) -> reqwest::RequestBuilder {
builder.bearer_auth("session-bearer")
}
impl ApiKeyProvider for AsyncProbe {
fn current_api_key(&self) -> Option<String> {
self.sync_calls.fetch_add(1, Ordering::SeqCst);
Some("sync-stale".into())
}
fn current_api_key_async(
&self,
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Option<String>> + Send + '_>>
{
let counter = self.async_calls.clone();
Box::pin(async move {
counter.fetch_add(1, Ordering::SeqCst);
Some("async-fresh".into())
})
}
#[async_trait::async_trait]
impl kigi_auth::AuthCredentialProvider for ProbeCredentials {
fn snapshot(&self) -> kigi_auth::CredentialSnapshot {
kigi_auth::CredentialSnapshot {
token: Some("session-bearer".into()),
..Default::default()
}
}
async fn refresh_after_unauthorized(&self) -> bool {
false
}
}
let sync_calls = Arc::new(AtomicU32::new(0));
let async_calls = Arc::new(AtomicU32::new(0));
let probe: kigi_tools::types::SharedApiKeyProvider = Arc::new(AsyncProbe {
sync_calls: sync_calls.clone(),
async_calls: async_calls.clone(),
});
const SESSION_ENDPOINT: &str = "https://api.kimi.com/coding/v1";
const FOREIGN_ENDPOINT: &str = "https://api.anthropic.com/v1";
let params = MemoryBackendParams {
fn params_at(base_url: &str, credentials: EndpointScopedCredentials) -> MemoryBackendParams {
MemoryBackendParams {
session_id: "s1".into(),
embed_config: Some(MemoryEmbeddingConfig {
model: Some("test-embed-model".into()),
..Default::default()
}),
embed_base_url: "http://example/v1".into(),
embed_api_key: Some("static-fallback".into()),
embed_base_url: base_url.into(),
embed_api_key: None,
search_config: MemorySearchConfig::default(),
watcher: None,
stale_claim_secs: 60,
search_source: "tool",
api_key_provider: Some(probe),
// No auth_credentials — forces the api_key_provider fallback path.
auth_credentials: None,
};
embedding_credentials: credentials,
}
}
let provider = params.make_embedding_provider().await;
/// The session bearer must never ride to a third-party embedding host.
///
/// `embed_base_url` is the CURRENT MODEL's endpoint, so a session on a BYOK
/// or subscription-OAuth model aims memory embeddings at that provider —
/// and the auth middleware stamps `Authorization` unconditionally.
#[tokio::test]
async fn session_credentials_are_withheld_from_a_foreign_endpoint() {
let scoped = EndpointScopedCredentials::for_endpoint(
FOREIGN_ENDPOINT,
false,
Some(Arc::new(ProbeCredentials)),
);
assert!(
scoped.is_empty(),
"an untrusted endpoint must drop both handles"
);
let provider = params_at(FOREIGN_ENDPOINT, scoped)
.make_embedding_provider()
.await;
assert!(
provider.is_none(),
"no credential may ride to a foreign endpoint, and there is no static key to fall back to"
);
}
/// The trusted-endpoint path still builds a credentialed provider.
#[tokio::test]
async fn session_credentials_ride_their_own_endpoint() {
let scoped = EndpointScopedCredentials::for_endpoint(
SESSION_ENDPOINT,
true,
Some(Arc::new(ProbeCredentials)),
);
assert!(
!scoped.is_empty(),
"a trusted endpoint must retain the handles"
);
let provider = params_at(SESSION_ENDPOINT, scoped)
.make_embedding_provider()
.await;
assert!(
provider.is_some(),
"provider must be built when model is set"
"the session endpoint keeps its refresh-capable credential"
);
assert_eq!(
async_calls.load(Ordering::SeqCst),
1,
"must call current_api_key_async exactly once per provider build"
}
/// The runtime re-check, not construction alone, is what guards the wire:
/// `MemoryBackendParams` is `Clone` and callers rewrite fields on the copy.
#[tokio::test]
async fn a_cloned_param_set_cannot_redirect_scoped_credentials() {
let scoped = EndpointScopedCredentials::for_endpoint(
SESSION_ENDPOINT,
true,
Some(Arc::new(ProbeCredentials)),
);
assert_eq!(
sync_calls.load(Ordering::SeqCst),
0,
"sync current_api_key must NOT be called — the async path is the contract"
let redirected = MemoryBackendParams {
embed_base_url: FOREIGN_ENDPOINT.into(),
..params_at(SESSION_ENDPOINT, scoped)
};
assert!(
redirected.make_embedding_provider().await.is_none(),
"credentials scoped to one endpoint must not follow a rewritten base_url"
);
}
/// A platform's own API key is not a session credential: it is resolved for
/// that platform and must keep serving that platform's endpoint.
#[tokio::test]
async fn a_platform_api_key_still_serves_its_own_endpoint() {
let params = MemoryBackendParams {
embed_api_key: Some("platform-key".into()),
..params_at(FOREIGN_ENDPOINT, EndpointScopedCredentials::none())
};
assert!(
params.make_embedding_provider().await.is_some(),
"withholding the session credential must not disable BYOK embeddings"
);
}
}
+1 -1
View File
@@ -31,7 +31,7 @@ pub mod storage;
pub mod text_utils;
pub mod watcher;
pub use backend::{MemoryBackendImpl, MemoryBackendParams};
pub use backend::{EndpointScopedCredentials, MemoryBackendImpl, MemoryBackendParams};
pub use index::{MemoryIndex, init_sqlite_vec};
pub use storage::{MemoryScope, MemoryStorage};