Skip to content
Open
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
69 changes: 65 additions & 4 deletions gnmi_server/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@ import (
"google.golang.org/grpc/authz"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/credentials/tls/certprovider"
"google.golang.org/grpc/metadata"
"google.golang.org/grpc/peer"
"google.golang.org/grpc/reflection"
"google.golang.org/grpc/security/advancedtls"
Expand Down Expand Up @@ -950,7 +951,17 @@ func IsNativeOrigin(origin string) bool {
}

// Get implements the Get RPC in gNMI spec.
func (s *Server) Get(ctx context.Context, req *gnmipb.GetRequest) (*gnmipb.GetResponse, error) {
func (s *Server) Get(ctx context.Context, req *gnmipb.GetRequest) (resp *gnmipb.GetResponse, err error) {
// GNMI-AUDIT logging
start := time.Now()
auditUser := extractUser(ctx)
auditPeer := extractPeer(ctx)

log.Infof("[GNMI-AUDIT] GetRequest user=%s peer=%s prefix=%v paths=%v type=%v",
auditUser, auditPeer, req.GetPrefix(), req.GetPath(), req.GetType())

defer logGnmiAuditResponse("Get", auditUser, auditPeer, start, &err)

common_utils.IncCounter(common_utils.GNMI_GET)

if req.GetType() != gnmipb.GetRequest_ALL {
Expand All @@ -975,7 +986,8 @@ func (s *Server) Get(ctx context.Context, req *gnmipb.GetRequest) (*gnmipb.GetRe
req.Path = newPaths
}

if err := s.checkEncodingAndModel(req.GetEncoding(), req.GetUseModels()); err != nil {
err = s.checkEncodingAndModel(req.GetEncoding(), req.GetUseModels())
if err != nil {
common_utils.IncCounter(common_utils.GNMI_GET_FAIL)
return nil, status.Error(codes.Unimplemented, err.Error())
}
Expand All @@ -994,7 +1006,6 @@ func (s *Server) Get(ctx context.Context, req *gnmipb.GetRequest) (*gnmipb.GetRe
log.V(2).Infof("GetRequest paths: %v", paths)

var dc sdc.Client
var err error
// Handle OPERATIONAL target directly without SONiC routing
if target == "OPERATIONAL" {
return s.handleOperationalGet(ctx, req, paths, prefix)
Expand Down Expand Up @@ -1090,7 +1101,16 @@ func SaveOnSetEnabled() error {
// SaveOnSetDisabeld does nothing.
func saveOnSetDisabled() error { return nil }

func (s *Server) Set(ctx context.Context, req *gnmipb.SetRequest) (*gnmipb.SetResponse, error) {
func (s *Server) Set(ctx context.Context, req *gnmipb.SetRequest) (resp *gnmipb.SetResponse, retErr error) {
start := time.Now()
auditUser := extractUser(ctx)
auditPeer := extractPeer(ctx)

log.Infof("[GNMI-AUDIT] SetRequest user=%s peer=%s prefix=%v updates=%d replaces=%d deletes=%d",
auditUser, auditPeer, req.GetPrefix(), len(req.GetUpdate()), len(req.GetReplace()), len(req.GetDelete()))

defer logGnmiAuditResponse("Set", auditUser, auditPeer, start, &retErr)

e := s.ReqFromMaster(req, &s.masterEID)
if e != nil {
return nil, e
Expand Down Expand Up @@ -1282,6 +1302,47 @@ func (s *Server) Capabilities(ctx context.Context, req *gnmipb.CapabilityRequest
Extension: exts}, nil
}

func extractUser(ctx context.Context) string {
rc, _ := common_utils.GetContext(ctx)
if rc != nil && rc.Auth.User != "" {
return rc.Auth.User
}

if md, ok := metadata.FromIncomingContext(ctx); ok {
for _, key := range []string{"username", "user", "x-remote-user"} {
if values := md.Get(key); len(values) > 0 && values[0] != "" {
return values[0]
}
}
}

if username, err := getUsername(ctx); err == nil && username != "" {
return username
}

return "unknown"
}

func extractPeer(ctx context.Context) string {
if pr, ok := peer.FromContext(ctx); ok && pr.Addr != nil {
return pr.Addr.String()
}

return "unknown"
}

func logGnmiAuditResponse(method, user, peer string, start time.Time, errPtr *error) {
duration := time.Since(start)

if errPtr != nil && *errPtr != nil {
log.Errorf("[GNMI-AUDIT] %sResponse user=%s peer=%s status=FAIL err=%v duration=%v",
method, user, peer, *errPtr, duration)
} else {
log.Infof("[GNMI-AUDIT] %sResponse user=%s peer=%s status=OK duration=%v",
method, user, peer, duration)
}
}

// Obtain the user name as the last element of the SPIFFE ID.
func getUsername(ctx context.Context) (string, error) {
pr, ok := peer.FromContext(ctx)
Expand Down
119 changes: 119 additions & 0 deletions gnmi_server/server_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -7269,6 +7269,125 @@ func TestSrvAdvConfig(t *testing.T) {
}
}

// TestGnmiAuditLogging verifies that [GNMI-AUDIT] log entries are produced
// with the correct method, user, peer, status, and error for Get and Set RPCs.
func TestGnmiAuditLogging(t *testing.T) {
type auditCall struct {
method string
user string
peer string
status string
err error
}

tests := []struct {
desc string
setupServer func(t *testing.T) *Server
makeRequest func(t *testing.T, gClient pb.GNMIClient, ctx context.Context)
wantMethod string
wantStatus string
wantErrMsg string
}{
{
desc: "Get request logged as FAIL on auth failure",
setupServer: func(t *testing.T) *Server { return createAuthServer(t, 8081) },
makeRequest: func(t *testing.T, gClient pb.GNMIClient, ctx context.Context) {
gClient.Get(ctx, &pb.GetRequest{})
},
wantMethod: "Get",
wantStatus: "FAIL",
wantErrMsg: "Unauthenticated",
},
{
desc: "Set request logged as FAIL on read-only server",
setupServer: func(t *testing.T) *Server { return createReadServer(t, 8081) },
makeRequest: func(t *testing.T, gClient pb.GNMIClient, ctx context.Context) {
gClient.Set(ctx, &pb.SetRequest{})
},
wantMethod: "Set",
wantStatus: "FAIL",
wantErrMsg: "read-only",
},
{
desc: "Set request logged as FAIL on auth failure",
setupServer: func(t *testing.T) *Server { return createAuthServer(t, 8081) },
makeRequest: func(t *testing.T, gClient pb.GNMIClient, ctx context.Context) {
gClient.Set(ctx, &pb.SetRequest{})
},
wantMethod: "Set",
wantStatus: "FAIL",
wantErrMsg: "Unauthenticated",
},
}

for _, tt := range tests {
t.Run(tt.desc, func(t *testing.T) {
var mu sync.Mutex
var captured auditCall

mock := gomonkey.ApplyFunc(logGnmiAuditResponse,
func(method, user, peer string, start time.Time, errPtr *error) {
mu.Lock()
defer mu.Unlock()
captured.method = method
captured.user = user
captured.peer = peer
if errPtr != nil && *errPtr != nil {
captured.status = "FAIL"
captured.err = *errPtr
} else {
captured.status = "OK"
}
})
defer mock.Reset()

s := tt.setupServer(t)
go runServer(t, s)
defer s.Stop()

tlsConfig := &tls.Config{InsecureSkipVerify: true}
opts := []grpc.DialOption{grpc.WithTransportCredentials(credentials.NewTLS(tlsConfig))}
conn, err := grpc.Dial("127.0.0.1:8081", opts...)
if err != nil {
t.Fatalf("Dialing failed: %v", err)
}
defer conn.Close()

gClient := pb.NewGNMIClient(conn)
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()

tt.makeRequest(t, gClient, ctx)
// Allow time for the deferred audit log to fire after RPC returns.
time.Sleep(200 * time.Millisecond)

mu.Lock()
gotMethod := captured.method
gotStatus := captured.status
gotErr := captured.err
gotPeer := captured.peer
mu.Unlock()

if gotMethod != tt.wantMethod {
t.Errorf("audit method: got %q, want %q", gotMethod, tt.wantMethod)
}
if gotStatus != tt.wantStatus {
t.Errorf("audit status: got %q, want %q", gotStatus, tt.wantStatus)
}
if gotPeer == "" || gotPeer == "unknown" {
t.Errorf("audit peer should be populated, got %q", gotPeer)
}
if tt.wantErrMsg != "" {
if gotErr == nil {
t.Errorf("expected audit error containing %q but got nil", tt.wantErrMsg)
} else if !strings.Contains(gotErr.Error(), tt.wantErrMsg) {
t.Errorf("audit error %q does not contain %q", gotErr.Error(), tt.wantErrMsg)
}
}
})
}
}

func TestSrvTestConfigLogsAndReturns(t *testing.T) {
// 1. Force error in the FIRST WriteFile (Server Certificate)
t.Run("CoverageCertWriteError", func(t *testing.T) {
Expand Down
Loading