fix(coding-agent): resolve models.json auth per request closes #1835
This commit is contained in:
@@ -24,7 +24,12 @@ import { existsSync, readFileSync } from "fs";
|
||||
import { join } from "path";
|
||||
import { getAgentDir } from "../config.js";
|
||||
import type { AuthStorage } from "./auth-storage.js";
|
||||
import { clearConfigValueCache, resolveConfigValue, resolveHeaders } from "./resolve-config-value.js";
|
||||
import {
|
||||
clearConfigValueCache,
|
||||
resolveConfigValueOrThrow,
|
||||
resolveConfigValueUncached,
|
||||
resolveHeadersOrThrow,
|
||||
} from "./resolve-config-value.js";
|
||||
|
||||
const Ajv = (AjvModule as any).default || AjvModule;
|
||||
const ajv = new Ajv();
|
||||
@@ -143,14 +148,29 @@ ajv.addSchema(ModelsConfigSchema, "ModelsConfig");
|
||||
|
||||
type ModelsConfig = Static<typeof ModelsConfigSchema>;
|
||||
|
||||
/** Provider override config (baseUrl, headers, apiKey, compat) without custom models */
|
||||
/** Provider override config (baseUrl, compat) without request auth/headers */
|
||||
interface ProviderOverride {
|
||||
baseUrl?: string;
|
||||
headers?: Record<string, string>;
|
||||
apiKey?: string;
|
||||
compat?: Model<Api>["compat"];
|
||||
}
|
||||
|
||||
interface ProviderRequestConfig {
|
||||
apiKey?: string;
|
||||
headers?: Record<string, string>;
|
||||
authHeader?: boolean;
|
||||
}
|
||||
|
||||
export type ResolvedRequestAuth =
|
||||
| {
|
||||
ok: true;
|
||||
apiKey?: string;
|
||||
headers?: Record<string, string>;
|
||||
}
|
||||
| {
|
||||
ok: false;
|
||||
error: string;
|
||||
};
|
||||
|
||||
/** Result of loading custom models from models.json */
|
||||
interface CustomModelsResult {
|
||||
models: Model<Api>[];
|
||||
@@ -220,12 +240,6 @@ function applyModelOverride(model: Model<Api>, override: ModelOverride): Model<A
|
||||
};
|
||||
}
|
||||
|
||||
// Merge headers
|
||||
if (override.headers) {
|
||||
const resolvedHeaders = resolveHeaders(override.headers);
|
||||
result.headers = resolvedHeaders ? { ...model.headers, ...resolvedHeaders } : model.headers;
|
||||
}
|
||||
|
||||
// Deep merge compat
|
||||
result.compat = mergeCompat(model.compat, override.compat);
|
||||
|
||||
@@ -240,7 +254,8 @@ export const clearApiKeyCache = clearConfigValueCache;
|
||||
*/
|
||||
export class ModelRegistry {
|
||||
private models: Model<Api>[] = [];
|
||||
private customProviderApiKeys: Map<string, string> = new Map();
|
||||
private providerRequestConfigs: Map<string, ProviderRequestConfig> = new Map();
|
||||
private modelRequestHeaders: Map<string, Record<string, string>> = new Map();
|
||||
private registeredProviders: Map<string, ProviderConfigInput> = new Map();
|
||||
private loadError: string | undefined = undefined;
|
||||
|
||||
@@ -248,16 +263,6 @@ export class ModelRegistry {
|
||||
readonly authStorage: AuthStorage,
|
||||
private modelsJsonPath: string | undefined = join(getAgentDir(), "models.json"),
|
||||
) {
|
||||
// Set up fallback resolver for custom provider API keys
|
||||
this.authStorage.setFallbackResolver((provider) => {
|
||||
const keyConfig = this.customProviderApiKeys.get(provider);
|
||||
if (keyConfig) {
|
||||
return resolveConfigValue(keyConfig);
|
||||
}
|
||||
return undefined;
|
||||
});
|
||||
|
||||
// Load models
|
||||
this.loadModels();
|
||||
}
|
||||
|
||||
@@ -265,7 +270,8 @@ export class ModelRegistry {
|
||||
* Reload models from disk (built-in + custom from models.json).
|
||||
*/
|
||||
refresh(): void {
|
||||
this.customProviderApiKeys.clear();
|
||||
this.providerRequestConfigs.clear();
|
||||
this.modelRequestHeaders.clear();
|
||||
this.loadError = undefined;
|
||||
|
||||
// Ensure dynamic API/OAuth registrations are rebuilt from current provider state.
|
||||
@@ -329,11 +335,9 @@ export class ModelRegistry {
|
||||
|
||||
// Apply provider-level baseUrl/headers/compat override
|
||||
if (providerOverride) {
|
||||
const resolvedHeaders = resolveHeaders(providerOverride.headers);
|
||||
model = {
|
||||
...model,
|
||||
baseUrl: providerOverride.baseUrl ?? model.baseUrl,
|
||||
headers: resolvedHeaders ? { ...model.headers, ...resolvedHeaders } : model.headers,
|
||||
compat: mergeCompat(model.compat, providerOverride.compat),
|
||||
};
|
||||
}
|
||||
@@ -388,23 +392,20 @@ export class ModelRegistry {
|
||||
const modelOverrides = new Map<string, Map<string, ModelOverride>>();
|
||||
|
||||
for (const [providerName, providerConfig] of Object.entries(config.providers)) {
|
||||
// Apply provider-level baseUrl/headers/apiKey/compat override to built-in models when configured.
|
||||
if (providerConfig.baseUrl || providerConfig.headers || providerConfig.apiKey || providerConfig.compat) {
|
||||
if (providerConfig.baseUrl || providerConfig.compat) {
|
||||
overrides.set(providerName, {
|
||||
baseUrl: providerConfig.baseUrl,
|
||||
headers: providerConfig.headers,
|
||||
apiKey: providerConfig.apiKey,
|
||||
compat: providerConfig.compat,
|
||||
});
|
||||
}
|
||||
|
||||
// Store API key for fallback resolver.
|
||||
if (providerConfig.apiKey) {
|
||||
this.customProviderApiKeys.set(providerName, providerConfig.apiKey);
|
||||
}
|
||||
this.storeProviderRequestConfig(providerName, providerConfig);
|
||||
|
||||
if (providerConfig.modelOverrides) {
|
||||
modelOverrides.set(providerName, new Map(Object.entries(providerConfig.modelOverrides)));
|
||||
for (const [modelId, modelOverride] of Object.entries(providerConfig.modelOverrides)) {
|
||||
this.storeModelHeaders(providerName, modelId, modelOverride.headers);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -469,29 +470,12 @@ export class ModelRegistry {
|
||||
const modelDefs = providerConfig.models ?? [];
|
||||
if (modelDefs.length === 0) continue; // Override-only, no custom models
|
||||
|
||||
// Store API key config for fallback resolver
|
||||
if (providerConfig.apiKey) {
|
||||
this.customProviderApiKeys.set(providerName, providerConfig.apiKey);
|
||||
}
|
||||
|
||||
for (const modelDef of modelDefs) {
|
||||
const api = modelDef.api || providerConfig.api;
|
||||
if (!api) continue;
|
||||
|
||||
// Merge headers: provider headers are base, model headers override
|
||||
// Resolve env vars and shell commands in header values
|
||||
const providerHeaders = resolveHeaders(providerConfig.headers);
|
||||
const modelHeaders = resolveHeaders(modelDef.headers);
|
||||
const compat = mergeCompat(providerConfig.compat, modelDef.compat);
|
||||
let headers = providerHeaders || modelHeaders ? { ...providerHeaders, ...modelHeaders } : undefined;
|
||||
|
||||
// If authHeader is true, add Authorization header with resolved API key
|
||||
if (providerConfig.authHeader && providerConfig.apiKey) {
|
||||
const resolvedKey = resolveConfigValue(providerConfig.apiKey);
|
||||
if (resolvedKey) {
|
||||
headers = { ...headers, Authorization: `Bearer ${resolvedKey}` };
|
||||
}
|
||||
}
|
||||
this.storeModelHeaders(providerName, modelDef.id, modelDef.headers);
|
||||
|
||||
// Provider baseUrl is required when custom models are defined.
|
||||
// Individual models can override it with modelDef.baseUrl.
|
||||
@@ -507,7 +491,7 @@ export class ModelRegistry {
|
||||
cost: modelDef.cost ?? defaultCost,
|
||||
contextWindow: modelDef.contextWindow ?? 128000,
|
||||
maxTokens: modelDef.maxTokens ?? 16384,
|
||||
headers,
|
||||
headers: undefined,
|
||||
compat,
|
||||
} as Model<Api>);
|
||||
}
|
||||
@@ -529,7 +513,7 @@ export class ModelRegistry {
|
||||
* This is a fast check that doesn't refresh OAuth tokens.
|
||||
*/
|
||||
getAvailable(): Model<Api>[] {
|
||||
return this.models.filter((m) => this.authStorage.hasAuth(m.provider));
|
||||
return this.models.filter((m) => this.hasConfiguredAuth(m));
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -542,15 +526,100 @@ export class ModelRegistry {
|
||||
/**
|
||||
* Get API key for a model.
|
||||
*/
|
||||
async getApiKey(model: Model<Api>): Promise<string | undefined> {
|
||||
return this.authStorage.getApiKey(model.provider);
|
||||
hasConfiguredAuth(model: Model<Api>): boolean {
|
||||
return (
|
||||
this.authStorage.hasAuth(model.provider) ||
|
||||
this.providerRequestConfigs.get(model.provider)?.apiKey !== undefined
|
||||
);
|
||||
}
|
||||
|
||||
private getModelRequestKey(provider: string, modelId: string): string {
|
||||
return `${provider}:${modelId}`;
|
||||
}
|
||||
|
||||
private storeProviderRequestConfig(
|
||||
providerName: string,
|
||||
config: {
|
||||
apiKey?: string;
|
||||
headers?: Record<string, string>;
|
||||
authHeader?: boolean;
|
||||
},
|
||||
): void {
|
||||
if (!config.apiKey && !config.headers && !config.authHeader) {
|
||||
return;
|
||||
}
|
||||
|
||||
this.providerRequestConfigs.set(providerName, {
|
||||
apiKey: config.apiKey,
|
||||
headers: config.headers,
|
||||
authHeader: config.authHeader,
|
||||
});
|
||||
}
|
||||
|
||||
private storeModelHeaders(providerName: string, modelId: string, headers?: Record<string, string>): void {
|
||||
const key = this.getModelRequestKey(providerName, modelId);
|
||||
if (!headers || Object.keys(headers).length === 0) {
|
||||
this.modelRequestHeaders.delete(key);
|
||||
return;
|
||||
}
|
||||
this.modelRequestHeaders.set(key, headers);
|
||||
}
|
||||
|
||||
/**
|
||||
* Get API key and request headers for a model.
|
||||
*/
|
||||
async getApiKeyAndHeaders(model: Model<Api>): Promise<ResolvedRequestAuth> {
|
||||
try {
|
||||
const providerConfig = this.providerRequestConfigs.get(model.provider);
|
||||
const apiKeyFromAuthStorage = await this.authStorage.getApiKey(model.provider, { includeFallback: false });
|
||||
const apiKey =
|
||||
apiKeyFromAuthStorage ??
|
||||
(providerConfig?.apiKey
|
||||
? resolveConfigValueOrThrow(providerConfig.apiKey, `API key for provider "${model.provider}"`)
|
||||
: undefined);
|
||||
|
||||
const providerHeaders = resolveHeadersOrThrow(providerConfig?.headers, `provider "${model.provider}"`);
|
||||
const modelHeaders = resolveHeadersOrThrow(
|
||||
this.modelRequestHeaders.get(this.getModelRequestKey(model.provider, model.id)),
|
||||
`model "${model.provider}/${model.id}"`,
|
||||
);
|
||||
|
||||
let headers =
|
||||
model.headers || providerHeaders || modelHeaders
|
||||
? { ...model.headers, ...providerHeaders, ...modelHeaders }
|
||||
: undefined;
|
||||
|
||||
if (providerConfig?.authHeader) {
|
||||
if (!apiKey) {
|
||||
return { ok: false, error: `No API key found for "${model.provider}"` };
|
||||
}
|
||||
headers = { ...headers, Authorization: `Bearer ${apiKey}` };
|
||||
}
|
||||
|
||||
return {
|
||||
ok: true,
|
||||
apiKey,
|
||||
headers: headers && Object.keys(headers).length > 0 ? headers : undefined,
|
||||
};
|
||||
} catch (error) {
|
||||
return {
|
||||
ok: false,
|
||||
error: error instanceof Error ? error.message : String(error),
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Get API key for a provider.
|
||||
*/
|
||||
async getApiKeyForProvider(provider: string): Promise<string | undefined> {
|
||||
return this.authStorage.getApiKey(provider);
|
||||
const apiKey = await this.authStorage.getApiKey(provider, { includeFallback: false });
|
||||
if (apiKey !== undefined) {
|
||||
return apiKey;
|
||||
}
|
||||
|
||||
const providerApiKey = this.providerRequestConfigs.get(provider)?.apiKey;
|
||||
return providerApiKey ? resolveConfigValueUncached(providerApiKey) : undefined;
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -586,7 +655,6 @@ export class ModelRegistry {
|
||||
unregisterProvider(providerName: string): void {
|
||||
if (!this.registeredProviders.has(providerName)) return;
|
||||
this.registeredProviders.delete(providerName);
|
||||
this.customProviderApiKeys.delete(providerName);
|
||||
this.refresh();
|
||||
}
|
||||
|
||||
@@ -637,10 +705,7 @@ export class ModelRegistry {
|
||||
);
|
||||
}
|
||||
|
||||
// Store API key for auth resolution
|
||||
if (config.apiKey) {
|
||||
this.customProviderApiKeys.set(providerName, config.apiKey);
|
||||
}
|
||||
this.storeProviderRequestConfig(providerName, config);
|
||||
|
||||
if (config.models && config.models.length > 0) {
|
||||
// Full replacement: remove existing models for this provider
|
||||
@@ -649,19 +714,7 @@ export class ModelRegistry {
|
||||
// Parse and add new models
|
||||
for (const modelDef of config.models) {
|
||||
const api = modelDef.api || config.api;
|
||||
|
||||
// Merge headers
|
||||
const providerHeaders = resolveHeaders(config.headers);
|
||||
const modelHeaders = resolveHeaders(modelDef.headers);
|
||||
let headers = providerHeaders || modelHeaders ? { ...providerHeaders, ...modelHeaders } : undefined;
|
||||
|
||||
// If authHeader is true, add Authorization header
|
||||
if (config.authHeader && config.apiKey) {
|
||||
const resolvedKey = resolveConfigValue(config.apiKey);
|
||||
if (resolvedKey) {
|
||||
headers = { ...headers, Authorization: `Bearer ${resolvedKey}` };
|
||||
}
|
||||
}
|
||||
this.storeModelHeaders(providerName, modelDef.id, modelDef.headers);
|
||||
|
||||
this.models.push({
|
||||
id: modelDef.id,
|
||||
@@ -674,7 +727,7 @@ export class ModelRegistry {
|
||||
cost: modelDef.cost,
|
||||
contextWindow: modelDef.contextWindow,
|
||||
maxTokens: modelDef.maxTokens,
|
||||
headers,
|
||||
headers: undefined,
|
||||
compat: modelDef.compat,
|
||||
} as Model<Api>);
|
||||
}
|
||||
@@ -686,15 +739,13 @@ export class ModelRegistry {
|
||||
this.models = config.oauth.modifyModels(this.models, cred);
|
||||
}
|
||||
}
|
||||
} else if (config.baseUrl) {
|
||||
// Override-only: update baseUrl/headers for existing models
|
||||
const resolvedHeaders = resolveHeaders(config.headers);
|
||||
} else if (config.baseUrl || config.headers) {
|
||||
// Override-only: update baseUrl for existing models. Request headers are resolved per request.
|
||||
this.models = this.models.map((m) => {
|
||||
if (m.provider !== providerName) return m;
|
||||
return {
|
||||
...m,
|
||||
baseUrl: config.baseUrl ?? m.baseUrl,
|
||||
headers: resolvedHeaders ? { ...m.headers, ...resolvedHeaders } : m.headers,
|
||||
};
|
||||
});
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user