diff --git a/gnmi_server/server.go b/gnmi_server/server.go index 55af5efc1..fa9502f66 100644 --- a/gnmi_server/server.go +++ b/gnmi_server/server.go @@ -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" @@ -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 { @@ -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()) } @@ -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) @@ -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 @@ -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) diff --git a/gnmi_server/server_test.go b/gnmi_server/server_test.go index a054958b7..e2141806e 100644 --- a/gnmi_server/server_test.go +++ b/gnmi_server/server_test.go @@ -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) {