feat: require bearer auth for accounts-backed ACP access
This commit is contained in:
parent
9bf00b3b4c
commit
3f41409363
@ -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,
|
||||
|
||||
@ -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")
|
||||
|
||||
@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@ -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
|
||||
}
|
||||
|
||||
@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@ -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
|
||||
}
|
||||
|
||||
@ -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")
|
||||
}
|
||||
}
|
||||
|
||||
Loading…
Reference in New Issue
Block a user