This is an automated email from the ASF dual-hosted git repository.

lahirujayathilake pushed a commit to branch auth-endpoints
in repository https://gitbox.apache.org/repos/asf/airavata-custos.git

commit a76bb3434951239e68f97770baa9620a92903508
Author: lahiruj <[email protected]>
AuthorDate: Tue Jun 16 19:25:42 2026 -0400

    Inject caller via context in integration test harness
---
 internal/server/integration_common_test.go    | 12 +++++++++++-
 internal/server/privilege_integration_test.go | 24 ++++++++++++------------
 2 files changed, 23 insertions(+), 13 deletions(-)

diff --git a/internal/server/integration_common_test.go 
b/internal/server/integration_common_test.go
index f3a9282c7..269c7dd67 100644
--- a/internal/server/integration_common_test.go
+++ b/internal/server/integration_common_test.go
@@ -20,6 +20,7 @@
 package server
 
 import (
+       "net/http"
        "os"
        "sync"
        "testing"
@@ -29,10 +30,19 @@ import (
 
        "github.com/apache/airavata-custos/internal/db"
        "github.com/apache/airavata-custos/pkg/events"
+       "github.com/apache/airavata-custos/pkg/identity"
        "github.com/apache/airavata-custos/pkg/models"
        "github.com/apache/airavata-custos/pkg/service"
 )
 
+// asCaller returns a copy of req with the verified caller attached to its
+// context. Integration tests construct requests against the server's mux
+// directly and rely on this helper to stand in for the auth middleware that
+// would attach identity in production.
+func asCaller(req *http.Request, userID string) *http.Request {
+       return req.WithContext(identity.WithCaller(req.Context(), 
&identity.Caller{UserID: userID}))
+}
+
 var (
        sharedDB     *sqlx.DB
        sharedDBOnce sync.Once
@@ -69,7 +79,7 @@ func setupTestStack(t *testing.T) (*sqlx.DB, 
*service.Service, *Server) {
        }
        truncateAll(t, sharedDB)
        svc := service.New(sharedDB, events.New())
-       return sharedDB, svc, New(svc)
+       return sharedDB, svc, New(svc, nil)
 }
 
 func truncateAll(t *testing.T, database *sqlx.DB) {
diff --git a/internal/server/privilege_integration_test.go 
b/internal/server/privilege_integration_test.go
index 18cc29d0f..7aa2b9948 100644
--- a/internal/server/privilege_integration_test.go
+++ b/internal/server/privilege_integration_test.go
@@ -44,7 +44,7 @@ func TestGetCallerPrivileges_NoGrants_ReturnsEmpty(t 
*testing.T) {
        user := seedUser(t, database, "[email protected]")
        rr := httptest.NewRecorder()
        req := httptest.NewRequest(http.MethodGet, "/user/privileges", nil)
-       req.Header.Set(callerHeader, user)
+       req = asCaller(req, user)
        srv.ServeHTTP(rr, req)
        if rr.Code != http.StatusOK {
                t.Fatalf("status: got %d, want 200", rr.Code)
@@ -67,7 +67,7 @@ func TestGetCallerPrivileges_WithGrants(t *testing.T) {
 
        rr := httptest.NewRecorder()
        req := httptest.NewRequest(http.MethodGet, "/user/privileges", nil)
-       req.Header.Set(callerHeader, user)
+       req = asCaller(req, user)
        srv.ServeHTTP(rr, req)
        if rr.Code != http.StatusOK {
                t.Fatalf("status: got %d, want 200", rr.Code)
@@ -88,7 +88,7 @@ func TestRequirePrivilege_NoGrants_403(t *testing.T) {
        user := seedUser(t, database, "[email protected]")
        rr := httptest.NewRecorder()
        req := httptest.NewRequest(http.MethodGet, "/privileges/catalog", nil)
-       req.Header.Set(callerHeader, user)
+       req = asCaller(req, user)
        srv.ServeHTTP(rr, req)
        if rr.Code != http.StatusForbidden {
                t.Errorf("status: got %d, want 403", rr.Code)
@@ -101,7 +101,7 @@ func TestRequirePrivilege_WithGrant_200(t *testing.T) {
        seedPrivilegeGrant(t, database, user)
        rr := httptest.NewRecorder()
        req := httptest.NewRequest(http.MethodGet, "/privileges/catalog", nil)
-       req.Header.Set(callerHeader, user)
+       req = asCaller(req, user)
        srv.ServeHTTP(rr, req)
        if rr.Code != http.StatusOK {
                t.Errorf("status: got %d, want 200", rr.Code)
@@ -124,7 +124,7 @@ func TestGrantPrivilegeEndpoint_HappyPath(t *testing.T) {
        body, _ := json.Marshal(map[string]any{"privilege": "amie:read", 
"reason": "ops view"})
        rr := httptest.NewRecorder()
        req := httptest.NewRequest(http.MethodPost, 
"/users/"+target+"/privileges", bytes.NewReader(body))
-       req.Header.Set(callerHeader, granter)
+       req = asCaller(req, granter)
        req.Header.Set("Content-Type", "application/json")
        srv.ServeHTTP(rr, req)
        if rr.Code != http.StatusCreated {
@@ -143,7 +143,7 @@ func TestGrantPrivilegeEndpoint_GranterWithoutMeta_403(t 
*testing.T) {
        body, _ := json.Marshal(map[string]any{"privilege": "amie:read"})
        rr := httptest.NewRecorder()
        req := httptest.NewRequest(http.MethodPost, 
"/users/"+target+"/privileges", bytes.NewReader(body))
-       req.Header.Set(callerHeader, plain)
+       req = asCaller(req, plain)
        srv.ServeHTTP(rr, req)
        if rr.Code != http.StatusForbidden {
                t.Errorf("status: got %d, want 403 (granter lacks 
privileges:grant)", rr.Code)
@@ -161,7 +161,7 @@ func TestRevokePrivilegeEndpoint_HappyPath(t *testing.T) {
        body, _ := json.Marshal(map[string]any{"reason": "rotated"})
        rr := httptest.NewRecorder()
        req := httptest.NewRequest(http.MethodDelete, 
"/users/"+target+"/privileges/amie:read", bytes.NewReader(body))
-       req.Header.Set(callerHeader, granter)
+       req = asCaller(req, granter)
        srv.ServeHTTP(rr, req)
        if rr.Code != http.StatusNoContent {
                t.Fatalf("status: got %d, want 204, body=%s", rr.Code, 
rr.Body.String())
@@ -177,7 +177,7 @@ func TestRevokePrivilegeEndpoint_SelfRevokeMeta_400(t 
*testing.T) {
        seedPrivilegeGrant(t, database, user)
        rr := httptest.NewRecorder()
        req := httptest.NewRequest(http.MethodDelete, 
"/users/"+user+"/privileges/privileges:grant", nil)
-       req.Header.Set(callerHeader, user)
+       req = asCaller(req, user)
        srv.ServeHTTP(rr, req)
        if rr.Code != http.StatusBadRequest {
                t.Errorf("status: got %d, want 400 (self-revoke of meta)", 
rr.Code)
@@ -194,7 +194,7 @@ func TestListUserPrivilegesEndpoint(t *testing.T) {
        }
        rr := httptest.NewRecorder()
        req := httptest.NewRequest(http.MethodGet, 
"/users/"+target+"/privileges", nil)
-       req.Header.Set(callerHeader, granter)
+       req = asCaller(req, granter)
        srv.ServeHTTP(rr, req)
        if rr.Code != http.StatusOK {
                t.Fatalf("status: got %d, want 200, body=%s", rr.Code, 
rr.Body.String())
@@ -218,7 +218,7 @@ func TestListPrivilegeHoldersEndpoint(t *testing.T) {
        }
        rr := httptest.NewRecorder()
        req := httptest.NewRequest(http.MethodGet, 
"/privileges/amie:read/holders", nil)
-       req.Header.Set(callerHeader, granter)
+       req = asCaller(req, granter)
        srv.ServeHTTP(rr, req)
        if rr.Code != http.StatusOK {
                t.Fatalf("status: got %d, want 200, body=%s", rr.Code, 
rr.Body.String())
@@ -242,7 +242,7 @@ func 
TestRequirePrivilege_StaleCacheStillReturns403AfterRevoke(t *testing.T) {
        // b warms the cache.
        warm := httptest.NewRecorder()
        warmReq := httptest.NewRequest(http.MethodGet, "/privileges/catalog", 
nil)
-       warmReq.Header.Set(callerHeader, b)
+       warmReq = asCaller(warmReq, b)
        srv.ServeHTTP(warm, warmReq)
        if warm.Code != http.StatusOK {
                t.Fatalf("warm-up: got %d, want 200", warm.Code)
@@ -259,7 +259,7 @@ func 
TestRequirePrivilege_StaleCacheStillReturns403AfterRevoke(t *testing.T) {
        // b retries and now gets 403.
        again := httptest.NewRecorder()
        againReq := httptest.NewRequest(http.MethodGet, "/privileges/catalog", 
nil)
-       againReq.Header.Set(callerHeader, b)
+       againReq = asCaller(againReq, b)
        srv.ServeHTTP(again, againReq)
        if again.Code != http.StatusForbidden {
                t.Errorf("post-revoke status for b: got %d, want 403", 
again.Code)

Reply via email to