feat: require bearer auth for accounts-backed ACP access

This commit is contained in:
Haitao Pan 2026-04-09 13:27:20 +08:00
parent 9bf00b3b4c
commit 3f41409363
7 changed files with 128 additions and 5 deletions

View File

@ -19,6 +19,7 @@ const (
externalProviderEndpointKey = "externalProviderEndpoint"
externalProviderAuthorizationHeaderKey = "externalProviderAuthorizationHeader"
externalProviderLabelKey = "externalProviderLabel"
inboundAuthorizationHeaderKey = "bridgeAuthorizationHeader"
)
func buildResolvedExecutionParams(
@ -74,6 +75,17 @@ func injectResolvedExternalProviderParams(
return params
}
func injectInboundAuthorizationHeader(params map[string]any, authorization string) map[string]any {
if params == nil {
params = map[string]any{}
}
authorization = strings.TrimSpace(authorization)
if authorization != "" {
params[inboundAuthorizationHeaderKey] = authorization
}
return params
}
func (s *Server) runGateway(
ctx context.Context,
method string,
@ -169,6 +181,7 @@ func sanitizeExternalACPParams(method string, params map[string]any) map[string]
delete(next, externalProviderEndpointKey)
delete(next, externalProviderAuthorizationHeaderKey)
delete(next, externalProviderLabelKey)
delete(next, inboundAuthorizationHeaderKey)
// Gateway-only fields are irrelevant in ACP single-agent forwarding.
normalizedMethod := strings.TrimSpace(method)
if normalizedMethod == "session.start" || normalizedMethod == "session.message" {
@ -187,11 +200,21 @@ func externalProviderFromParams(params map[string]any) (syncedProvider, bool) {
ProviderID: strings.TrimSpace(shared.StringArg(params, "provider", "")),
Label: strings.TrimSpace(shared.StringArg(params, externalProviderLabelKey, "")),
Endpoint: endpoint,
AuthorizationHeader: strings.TrimSpace(shared.StringArg(params, externalProviderAuthorizationHeaderKey, "")),
AuthorizationHeader: fallbackAuthorizationHeader(
strings.TrimSpace(shared.StringArg(params, externalProviderAuthorizationHeaderKey, "")),
strings.TrimSpace(shared.StringArg(params, inboundAuthorizationHeaderKey, "")),
),
Enabled: true,
}, true
}
func fallbackAuthorizationHeader(explicit, inbound string) string {
if strings.TrimSpace(explicit) != "" {
return strings.TrimSpace(explicit)
}
return strings.TrimSpace(inbound)
}
func requestExternalACP(
ctx context.Context,
endpoint,

View File

@ -18,6 +18,7 @@ import (
"xworkmate-bridge/internal/gatewayruntime"
"xworkmate-bridge/internal/mounts"
"xworkmate-bridge/internal/router"
"xworkmate-bridge/internal/service"
"xworkmate-bridge/internal/shared"
)
@ -49,6 +50,7 @@ type Server struct {
queues map[string]chan task
gateway *gatewayruntime.Manager
providerCatalog map[string]syncedProvider
authService *service.StaticTokenAuthService
}
var wsUpgrader = websocket.Upgrader{
@ -99,6 +101,7 @@ func NewServer() *Server {
queues: make(map[string]chan task),
gateway: gatewayruntime.NewManager(),
providerCatalog: make(map[string]syncedProvider),
authService: service.NewStaticTokenAuthService(""),
}
}
@ -114,9 +117,19 @@ func (s *Server) HandleWebSocket(w http.ResponseWriter, r *http.Request) {
)
return
}
if !s.authorized(r) {
s.writeJSONError(
w,
nil,
http.StatusUnauthorized,
-32001,
"missing bearer authorization",
)
return
}
upgrader := wsUpgrader
upgrader.CheckOrigin = func(req *http.Request) bool {
return s.originAllowed(req.Header.Get("Origin"))
return s.originAllowed(req.Header.Get("Origin")) && s.authorized(req)
}
conn, err := upgrader.Upgrade(w, r, nil)
if err != nil {
@ -141,6 +154,10 @@ func (s *Server) HandleWebSocket(w http.ResponseWriter, r *http.Request) {
notify(shared.ErrorEnvelope(nil, -32700, err.Error()))
continue
}
request.Params = injectInboundAuthorizationHeader(
request.Params,
r.Header.Get("Authorization"),
)
response, rpcErr := s.handleRequest(request, notify)
if request.ID == nil {
continue
@ -180,6 +197,16 @@ func (s *Server) HandleRPC(w http.ResponseWriter, r *http.Request) {
)
return
}
if !s.authorized(r) {
s.writeJSONError(
w,
nil,
http.StatusUnauthorized,
-32001,
"missing bearer authorization",
)
return
}
payload, err := io.ReadAll(r.Body)
if err != nil {
s.writeJSONError(w, nil, http.StatusBadRequest, -32600, "invalid body")

View File

@ -45,6 +45,23 @@ func TestHandleRPCAllowsPreflightForConfiguredOrigin(t *testing.T) {
}
}
func TestHandleRPCRequiresBearerAuthorization(t *testing.T) {
server := NewServer()
recorder := httptest.NewRecorder()
request := httptest.NewRequest(
http.MethodPost,
"http://127.0.0.1/acp/rpc",
strings.NewReader(`{"jsonrpc":"2.0","id":1,"method":"acp.capabilities"}`),
)
request.Header.Set("Content-Type", "application/json")
server.HandleRPC(recorder, request)
if recorder.Code != http.StatusUnauthorized {
t.Fatalf("expected 401, got %d", recorder.Code)
}
}
func TestHandleRPCRejectsUnknownOrigin(t *testing.T) {
t.Setenv("ACP_ALLOWED_ORIGINS", "https://xworkmate.svc.plus")
@ -57,6 +74,7 @@ func TestHandleRPCRejectsUnknownOrigin(t *testing.T) {
)
request.Header.Set("Origin", "https://evil.example.com")
request.Header.Set("Content-Type", "application/json")
request.Header.Set("Authorization", "Bearer test")
server.HandleRPC(recorder, request)
@ -76,6 +94,7 @@ func TestHandleRPCMethodErrorUsesJSONEnvelope(t *testing.T) {
server := NewServer()
recorder := httptest.NewRecorder()
request := httptest.NewRequest(http.MethodGet, "http://127.0.0.1/acp/rpc", nil)
request.Header.Set("Authorization", "Bearer test")
server.HandleRPC(recorder, request)
@ -96,6 +115,7 @@ func TestHandleRPCCapabilitiesStillReturnsJSONResult(t *testing.T) {
strings.NewReader(`{"jsonrpc":"2.0","id":1,"method":"acp.capabilities"}`),
)
request.Header.Set("Content-Type", "application/json")
request.Header.Set("Authorization", "Bearer test")
server.HandleRPC(recorder, request)
@ -109,3 +129,18 @@ func TestHandleRPCCapabilitiesStillReturnsJSONResult(t *testing.T) {
t.Fatalf("expected capabilities response, got %q", recorder.Body.String())
}
}
func TestHandleWebSocketRequiresBearerAuthorization(t *testing.T) {
t.Setenv("ACP_ALLOWED_ORIGINS", "https://xworkmate.svc.plus")
server := NewServer()
recorder := httptest.NewRecorder()
request := httptest.NewRequest(http.MethodGet, "http://127.0.0.1/acp", nil)
request.Header.Set("Origin", "https://xworkmate.svc.plus")
server.HandleWebSocket(recorder, request)
if recorder.Code != http.StatusUnauthorized {
t.Fatalf("expected 401, got %d", recorder.Code)
}
}

View File

@ -21,7 +21,7 @@ func (h *TokenAuthHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
return
}
token := r.Header.Get("Authorization")
if !h.service.ValidateToken(token) {
if !h.service.ValidateAuthorizationHeader(token) {
http.Error(w, "unauthorized", http.StatusUnauthorized)
return
}

View File

@ -20,3 +20,15 @@ func TestTokenAuthHandlerServeHTTP(t *testing.T) {
t.Fatalf("expected 200, got %d", rec.Code)
}
}
func TestTokenAuthHandlerRejectsMissingBearer(t *testing.T) {
h := NewTokenAuthHandler(service.NewStaticTokenAuthService(""))
req := httptest.NewRequest(http.MethodGet, "/", nil)
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusUnauthorized {
t.Fatalf("expected 401, got %d", rec.Code)
}
}

View File

@ -13,6 +13,19 @@ func NewStaticTokenAuthService(expectedToken string) *StaticTokenAuthService {
}
func (s *StaticTokenAuthService) ValidateToken(token string) bool {
token = strings.TrimSpace(token)
return token != "" && token == s.expectedToken
return s.ValidateAuthorizationHeader(token)
}
func (s *StaticTokenAuthService) ValidateAuthorizationHeader(header string) bool {
header = strings.TrimSpace(header)
if header == "" {
return false
}
if s.expectedToken == "" {
if !strings.HasPrefix(strings.ToLower(header), "bearer ") {
return false
}
return strings.TrimSpace(header[len("Bearer "):]) != ""
}
return header == s.expectedToken
}

View File

@ -11,3 +11,16 @@ func TestStaticTokenAuthServiceValidateToken(t *testing.T) {
t.Fatal("expected invalid token")
}
}
func TestStaticTokenAuthServiceValidateAuthorizationHeaderAsBearer(t *testing.T) {
svc := NewStaticTokenAuthService("")
if !svc.ValidateAuthorizationHeader("Bearer test-token") {
t.Fatal("expected bearer header to be accepted")
}
if svc.ValidateAuthorizationHeader("Basic abc") {
t.Fatal("expected non-bearer header to be rejected")
}
if svc.ValidateAuthorizationHeader("Bearer ") {
t.Fatal("expected empty bearer token to be rejected")
}
}