diff --git a/open-sse/services/tokenRefresh.ts b/open-sse/services/tokenRefresh.ts index 507e9ff27b..2528c167bc 100755 --- a/open-sse/services/tokenRefresh.ts +++ b/open-sse/services/tokenRefresh.ts @@ -50,9 +50,22 @@ export const REFRESH_LEAD_MS: Record = { /** * Get the proactive refresh lead time (ms) for a given provider. - * Falls back to TOKEN_EXPIRY_BUFFER_MS (5 min) when not explicitly listed. + * + * Precedence: + * 1. A per-connection override in `providerSpecificData.refreshLeadMs` + * (must be a positive finite number), so an operator can tune the lead + * time for a single connection without touching the provider defaults. + * 2. The provider default from REFRESH_LEAD_MS. + * 3. TOKEN_EXPIRY_BUFFER_MS (5 min) when nothing else applies. */ -export function getRefreshLeadMs(provider: string): number { +export function getRefreshLeadMs( + provider: string, + providerSpecificData?: { refreshLeadMs?: unknown } | null +): number { + const override = providerSpecificData?.refreshLeadMs; + if (typeof override === "number" && Number.isFinite(override) && override > 0) { + return override; + } return REFRESH_LEAD_MS[provider] ?? TOKEN_EXPIRY_BUFFER_MS; } diff --git a/src/sse/services/tokenRefresh.ts b/src/sse/services/tokenRefresh.ts index 63415b09b0..ab51be2742 100755 --- a/src/sse/services/tokenRefresh.ts +++ b/src/sse/services/tokenRefresh.ts @@ -174,7 +174,7 @@ export async function checkAndRefreshToken(provider: string, credentials: any) { if (updatedCredentials.expiresAt) { const expiresAt = new Date(updatedCredentials.expiresAt).getTime(); const now = Date.now(); - const refreshLead = _getRefreshLeadMs(provider); + const refreshLead = _getRefreshLeadMs(provider, updatedCredentials.providerSpecificData); if (expiresAt - now < refreshLead) { log.info("TOKEN_REFRESH", "Token expiring soon, refreshing proactively", { diff --git a/tests/unit/service-token-refresh.test.ts b/tests/unit/service-token-refresh.test.ts index f254d54e7d..42622907fe 100644 --- a/tests/unit/service-token-refresh.test.ts +++ b/tests/unit/service-token-refresh.test.ts @@ -17,6 +17,32 @@ describe("tokenRefresh helpers", () => { assert.equal(mod.getRefreshLeadMs("unknown-provider"), mod.TOKEN_EXPIRY_BUFFER_MS); assert.equal(mod.getRefreshLeadMs(""), mod.TOKEN_EXPIRY_BUFFER_MS); }); + + it("honors a positive per-connection refreshLeadMs override", () => { + // Override beats both the provider default and the fallback buffer. + assert.equal(mod.getRefreshLeadMs("codex", { refreshLeadMs: 90_000 }), 90_000); + assert.equal( + mod.getRefreshLeadMs("unknown-provider", { refreshLeadMs: 12_345 }), + 12_345 + ); + }); + + it("ignores invalid or non-positive override values", () => { + // Falls through to provider default / buffer when the override is unusable. + assert.equal(mod.getRefreshLeadMs("codex", null), 5 * 60 * 1000); + assert.equal(mod.getRefreshLeadMs("codex", {}), 5 * 60 * 1000); + assert.equal(mod.getRefreshLeadMs("codex", { refreshLeadMs: 0 }), 5 * 60 * 1000); + assert.equal(mod.getRefreshLeadMs("codex", { refreshLeadMs: -1 }), 5 * 60 * 1000); + assert.equal( + mod.getRefreshLeadMs("codex", { refreshLeadMs: "60000" as unknown as number }), + 5 * 60 * 1000 + ); + assert.equal(mod.getRefreshLeadMs("codex", { refreshLeadMs: NaN }), 5 * 60 * 1000); + assert.equal( + mod.getRefreshLeadMs("unknown-provider", { refreshLeadMs: -5 }), + mod.TOKEN_EXPIRY_BUFFER_MS + ); + }); }); describe("supportsTokenRefresh", () => {