From d815e6849056077aca2dc7580decb2fd7bde760a Mon Sep 17 00:00:00 2001 From: Shreyansh Sancheti <43677304+shreyanshjain7174@users.noreply.github.com> Date: Sat, 29 Aug 2026 22:11:55 +0530 Subject: [PATCH] feat: refresh expiring OAuth tokens before dispatch Signed-off-by: Shreyansh Sancheti <43677304+shreyanshjain7174@users.noreply.github.com> --- README.md | 1 + cmd/agentgw/main.go | 12 +++- internal/gateway/gateway.go | 21 +++++-- internal/gateway/gateway_test.go | 103 +++++++++++++++++++++++++++++++ internal/oauth/handler.go | 38 ++++++++++-- internal/oauth/refresh_test.go | 51 +++++++++++++++ 6 files changed, 214 insertions(+), 12 deletions(-) create mode 100644 internal/oauth/refresh_test.go diff --git a/README.md b/README.md index 08b0e56..58f106c 100644 --- a/README.md +++ b/README.md @@ -251,6 +251,7 @@ if sdk.IsRateLimited(err) { 4. **Signed receipts.** Every authenticated action attempt commits one Ed25519-signed, hash-chained receipt before the response is returned — verify offline with `agentgate-verify`, no gateway state or private key needed. 5. **Rate limiting.** Per-(agent, service) token bucket prevents runaway API usage. 6. **OAuth state encrypted** with AES-256-GCM and expires after 10 minutes. +7. **OAuth refresh.** Configured OAuth providers refresh tokens expiring within 5 minutes before dispatch. A failed refresh returns `token_expired` without calling the SaaS API. ## Configuration diff --git a/cmd/agentgw/main.go b/cmd/agentgw/main.go index fb06fcc..4a545e5 100644 --- a/cmd/agentgw/main.go +++ b/cmd/agentgw/main.go @@ -124,6 +124,7 @@ func main() { os.Exit(1) } delegationVerifier := delegation.NewVerifier(rootPub) + oauthProviders := buildOAuthProviders(reg, logger) srv := gateway.New(gateway.Config{ Registry: reg, @@ -133,6 +134,7 @@ func main() { Receipts: ledger, Limiter: limiter, Delegation: delegationVerifier, + Refreshers: buildTokenRefreshers(oauthProviders), }) adminSecret := os.Getenv("AGENTGATE_ADMIN_SECRET") @@ -145,7 +147,7 @@ func main() { publicURL = "http://localhost:8080" } - oauthHandler := oauth.NewCallbackHandler(buildOAuthProviders(reg, logger), vaultStore, masterKey, publicURL, logger) + oauthHandler := oauth.NewCallbackHandler(oauthProviders, vaultStore, masterKey, publicURL, logger) adminHandler := admin.NewHandler(keyStore, oauthHandler, vaultStore, adminSecret, logger) mux := http.NewServeMux() @@ -271,3 +273,11 @@ func buildOAuthProviders(reg *registry.Registry, logger *slog.Logger) map[string } return providers } + +func buildTokenRefreshers(providers map[string]*oauth.Provider) map[string]vault.RefreshFunc { + refreshers := make(map[string]vault.RefreshFunc, len(providers)) + for name, provider := range providers { + refreshers[name] = oauth.NewRefreshFunc(provider) + } + return refreshers +} diff --git a/internal/gateway/gateway.go b/internal/gateway/gateway.go index 6af60c1..874a4b1 100644 --- a/internal/gateway/gateway.go +++ b/internal/gateway/gateway.go @@ -77,10 +77,11 @@ type Config struct { Vault vault.Store HTTPClient *http.Client Logger *slog.Logger - Authorizer AgentAuthorizer // required: verifies API keys and scopes - Receipts ReceiptRecorder // required: commits one receipt per attempt - Limiter RequestLimiter // optional: nil disables rate limiting - Delegation DelegationVerifier // optional: nil rejects any request presenting a delegation token + Authorizer AgentAuthorizer // required: verifies API keys and scopes + Receipts ReceiptRecorder // required: commits one receipt per attempt + Limiter RequestLimiter // optional: nil disables rate limiting + Delegation DelegationVerifier // optional: nil rejects any request presenting a delegation token + Refreshers map[string]vault.RefreshFunc // optional: refreshes expiring OAuth tokens by service } // Server is the gateway HTTP server. @@ -399,8 +400,18 @@ func (s *Server) executeAttempt(ctx context.Context, att *attempt) *outcome { errorCode: "token_missing", } } + if refresher := s.cfg.Refreshers[att.req.Service]; refresher != nil { + tok, err = vault.GetOrRefresh(s.cfg.Vault, att.req.OnBehalfOf, att.req.Service, refresher) + if err != nil { + return &outcome{ + status: http.StatusForbidden, + body: ErrorResponse{Error: "token expired — user must re-authenticate", Code: "token_expired"}, + policyDecision: "allow", + errorCode: "token_expired", + } + } + } if tok.IsExpired() { - // TODO: auto-refresh using refresh_token return &outcome{ status: http.StatusForbidden, body: ErrorResponse{Error: "token expired — user must re-authenticate", Code: "token_expired"}, diff --git a/internal/gateway/gateway_test.go b/internal/gateway/gateway_test.go index e89ffcc..ddb7b38 100644 --- a/internal/gateway/gateway_test.go +++ b/internal/gateway/gateway_test.go @@ -414,6 +414,109 @@ services: } } +func TestAct_RefreshesExpiringTokenBeforeDispatch(t *testing.T) { + t.Parallel() + + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if got := r.Header.Get("Authorization"); got != "Bearer refreshed-token" { + t.Errorf("upstream authorization = %q, want refreshed token", got) + } + w.Header().Set("Content-Type", "application/json") + fmt.Fprint(w, `[]`) + })) + defer upstream.Close() + + reg := registry.New() + _ = reg.LoadBytes([]byte(fmt.Sprintf(` +services: + github: + base_url: %s + auth: + type: oauth2 + actions: + list_repos: + method: GET + path: /user/repos +`, upstream.URL))) + + key := make([]byte, 32) + store, _ := vault.NewMemoryStore(key) + _ = store.Put("user-1", "github", vault.Token{ + AccessToken: "expired-token", + RefreshToken: "refresh-token", + ExpiresAt: time.Now().Add(-time.Minute), + }) + + refreshed := false + srv := New(Config{ + Registry: reg, + Vault: store, + Authorizer: &fakeAuthorizer{keys: map[string]*auth.AgentKey{ + "test-key": {ID: "test-agent", AllowedServices: []string{"*"}, AllowedUsers: []string{"*"}}, + }}, + Receipts: &fakeReceipts{}, + Refreshers: map[string]vault.RefreshFunc{ + "github": func(refreshToken string) (string, string, time.Duration, error) { + if refreshToken != "refresh-token" { + t.Errorf("refresh token = %q, want stored refresh token", refreshToken) + } + refreshed = true + return "refreshed-token", "rotated-refresh-token", time.Hour, nil + }, + }, + }) + + req := httptest.NewRequest("POST", "/v1/act", strings.NewReader(`{"service":"github","action":"list_repos","on_behalf_of":"user-1"}`)) + req.Header.Set("Authorization", "Bearer test-key") + w := httptest.NewRecorder() + srv.ServeHTTP(w, req) + + if w.Code != http.StatusOK { + t.Fatalf("status = %d, want 200", w.Code) + } + if !refreshed { + t.Fatal("refresh function was not called") + } + stored, err := store.Get("user-1", "github") + if err != nil { + t.Fatalf("get refreshed token: %v", err) + } + if stored.AccessToken != "refreshed-token" || stored.RefreshToken != "rotated-refresh-token" { + t.Fatalf("stored token = %+v, want refreshed values", stored) + } +} + +func TestAct_RefreshFailureBlocksDispatchAndIsReceipted(t *testing.T) { + t.Parallel() + + srv, store, receipts := testSetup(t) + _ = store.Put("user-1", "github", vault.Token{ + AccessToken: "expired-token", + RefreshToken: "refresh-token", + ExpiresAt: time.Now().Add(-time.Minute), + }) + srv.cfg.Refreshers = map[string]vault.RefreshFunc{ + "github": func(string) (string, string, time.Duration, error) { + return "", "", 0, errors.New("provider rejected refresh token") + }, + } + + req := httptest.NewRequest("POST", "/v1/act", strings.NewReader(`{"service":"github","action":"list_repos","on_behalf_of":"user-1"}`)) + req.Header.Set("Authorization", "Bearer test-key") + w := httptest.NewRecorder() + srv.ServeHTTP(w, req) + + if w.Code != http.StatusForbidden { + t.Fatalf("status = %d, want 403", w.Code) + } + if receipts.count() != 1 { + t.Fatalf("receipts = %d, want 1", receipts.count()) + } + if got := receipts.last().Error; got != "token_expired" { + t.Fatalf("receipt error = %q, want token_expired", got) + } +} + func TestAct_PathParams(t *testing.T) { t.Parallel() diff --git a/internal/oauth/handler.go b/internal/oauth/handler.go index 83bdcb3..d8c0ad2 100644 --- a/internal/oauth/handler.go +++ b/internal/oauth/handler.go @@ -150,17 +150,42 @@ func (h *CallbackHandler) exchangeCode(provider *Provider, code, service string) "client_secret": {provider.ClientSecret}, "redirect_uri": {h.callbackBase + "/auth/callback/" + service}, } + tok, _, err := exchangeToken(provider, data) + return tok, err +} + +// NewRefreshFunc returns the refresh-token exchange for one configured OAuth +// provider. It returns an error when the provider omits an expiry because the +// caller must not persist a token that will immediately be treated as expired. +func NewRefreshFunc(provider *Provider) vault.RefreshFunc { + return func(refreshToken string) (string, string, time.Duration, error) { + tok, expiresIn, err := exchangeToken(provider, url.Values{ + "grant_type": {"refresh_token"}, + "refresh_token": {refreshToken}, + "client_id": {provider.ClientID}, + "client_secret": {provider.ClientSecret}, + }) + if err != nil { + return "", "", 0, err + } + if expiresIn <= 0 { + return "", "", 0, fmt.Errorf("oauth: refresh response has no expires_in") + } + return tok.AccessToken, tok.RefreshToken, expiresIn, nil + } +} +func exchangeToken(provider *Provider, data url.Values) (*vault.Token, time.Duration, error) { resp, err := http.PostForm(provider.TokenURL, data) if err != nil { - return nil, fmt.Errorf("oauth: exchange: %w", err) + return nil, 0, fmt.Errorf("oauth: exchange: %w", err) } defer resp.Body.Close() body, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20)) if resp.StatusCode != 200 { - return nil, fmt.Errorf("oauth: exchange returned %d: %s", resp.StatusCode, string(body)) + return nil, 0, fmt.Errorf("oauth: exchange returned %d: %s", resp.StatusCode, string(body)) } var tokenResp struct { @@ -171,7 +196,7 @@ func (h *CallbackHandler) exchangeCode(provider *Provider, code, service string) Scope string `json:"scope"` } if err := json.Unmarshal(body, &tokenResp); err != nil { - return nil, fmt.Errorf("oauth: parse token response: %w", err) + return nil, 0, fmt.Errorf("oauth: parse token response: %w", err) } tok := &vault.Token{ @@ -179,12 +204,13 @@ func (h *CallbackHandler) exchangeCode(provider *Provider, code, service string) RefreshToken: tokenResp.RefreshToken, TokenType: tokenResp.TokenType, } - if tokenResp.ExpiresIn > 0 { - tok.ExpiresAt = time.Now().Add(time.Duration(tokenResp.ExpiresIn) * time.Second) + expiresIn := time.Duration(tokenResp.ExpiresIn) * time.Second + if expiresIn > 0 { + tok.ExpiresAt = time.Now().Add(expiresIn) } if tokenResp.Scope != "" { tok.Scopes = strings.Split(tokenResp.Scope, " ") } - return tok, nil + return tok, expiresIn, nil } diff --git a/internal/oauth/refresh_test.go b/internal/oauth/refresh_test.go new file mode 100644 index 0000000..d2b0020 --- /dev/null +++ b/internal/oauth/refresh_test.go @@ -0,0 +1,51 @@ +package oauth + +import ( + "io" + "net/http" + "net/http/httptest" + "testing" + "time" +) + +func TestNewRefreshFuncExchangesAndRotatesToken(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if err := r.ParseForm(); err != nil { + t.Fatalf("parse form: %v", err) + } + if got := r.Form.Get("grant_type"); got != "refresh_token" { + t.Errorf("grant_type = %q, want refresh_token", got) + } + if got := r.Form.Get("refresh_token"); got != "old-refresh" { + t.Errorf("refresh_token = %q, want old-refresh", got) + } + if got := r.Form.Get("client_id"); got != "client-id" { + t.Errorf("client_id = %q, want client-id", got) + } + if got := r.Form.Get("client_secret"); got != "client-secret" { + t.Errorf("client_secret = %q, want client-secret", got) + } + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"access_token":"new-access","refresh_token":"new-refresh","token_type":"Bearer","expires_in":3600,"scope":"repo read:org"}`) + })) + defer server.Close() + + refresh := NewRefreshFunc(&Provider{ + ClientID: "client-id", + ClientSecret: "client-secret", + TokenURL: server.URL, + }) + + access, rotated, expiresIn, err := refresh("old-refresh") + if err != nil { + t.Fatalf("refresh: %v", err) + } + if access != "new-access" || rotated != "new-refresh" { + t.Fatalf("tokens = (%q, %q), want refreshed values", access, rotated) + } + if expiresIn != time.Hour { + t.Fatalf("expiresIn = %s, want %s", expiresIn, time.Hour) + } +}