Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
12 changes: 11 additions & 1 deletion cmd/agentgw/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -124,6 +124,7 @@ func main() {
os.Exit(1)
}
delegationVerifier := delegation.NewVerifier(rootPub)
oauthProviders := buildOAuthProviders(reg, logger)

srv := gateway.New(gateway.Config{
Registry: reg,
Expand All @@ -133,6 +134,7 @@ func main() {
Receipts: ledger,
Limiter: limiter,
Delegation: delegationVerifier,
Refreshers: buildTokenRefreshers(oauthProviders),
})

adminSecret := os.Getenv("AGENTGATE_ADMIN_SECRET")
Expand All @@ -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()
Expand Down Expand Up @@ -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
}
21 changes: 16 additions & 5 deletions internal/gateway/gateway.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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",
}
}
Comment on lines +403 to +412
}
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"},
Expand Down
103 changes: 103 additions & 0 deletions internal/gateway/gateway_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand Down
38 changes: 32 additions & 6 deletions internal/oauth/handler.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
Comment on lines +178 to 182
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 {
Expand All @@ -171,20 +196,21 @@ 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{
AccessToken: tokenResp.AccessToken,
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
}
51 changes: 51 additions & 0 deletions internal/oauth/refresh_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
}
Loading