From 612fa05dd3f7dbb41f6f9c995fdd821ff5848584 Mon Sep 17 00:00:00 2001 From: Tomas Touceda Date: Mon, 23 Aug 2021 19:40:00 -0300 Subject: [PATCH] Log errors when osquery endpoints have issues (#1764) --- .../issue-1758-log-errors-in-osquery-requests | 1 + server/service/endpoint_middleware.go | 32 +++++++- server/service/handler.go | 22 +++--- server/service/integration_test.go | 76 +++++++++++++++++++ 4 files changed, 119 insertions(+), 12 deletions(-) create mode 100644 changes/issue-1758-log-errors-in-osquery-requests diff --git a/changes/issue-1758-log-errors-in-osquery-requests b/changes/issue-1758-log-errors-in-osquery-requests new file mode 100644 index 0000000000..dcc2f7677d --- /dev/null +++ b/changes/issue-1758-log-errors-in-osquery-requests @@ -0,0 +1 @@ +* Log errors when osquery endpoints have issues. diff --git a/server/service/endpoint_middleware.go b/server/service/endpoint_middleware.go index 91f40fbb5e..b4d61ae896 100644 --- a/server/service/endpoint_middleware.go +++ b/server/service/endpoint_middleware.go @@ -25,11 +25,23 @@ func authenticatedHost(svc fleet.Service, next endpoint.Endpoint) endpoint.Endpo host, err := svc.AuthenticateHost(ctx, nodeKey) if err != nil { + logging.WithErr(ctx, err) return nil, err } ctx = hostctx.NewContext(ctx, *host) - return next(ctx, request) + resp, err := next(ctx, request) + if err != nil { + logging.WithErr(ctx, err) + return nil, err + } + if errResp, ok := resp.(errorer); ok { + err = errResp.error() + if err != nil { + logging.WithErr(ctx, err) + } + } + return resp, nil } } @@ -100,6 +112,24 @@ func authenticatedUser(svc fleet.Service, next endpoint.Endpoint) endpoint.Endpo } } +// logged wraps an endpoint and adds the error if the context supports it +func logged(next endpoint.Endpoint) endpoint.Endpoint { + return func(ctx context.Context, request interface{}) (response interface{}, err error) { + res, err := next(ctx, request) + if err != nil { + logging.WithErr(ctx, err) + return nil, err + } + if errResp, ok := res.(errorer); ok { + err = errResp.error() + if err != nil { + logging.WithErr(ctx, err) + } + } + return res, nil + } +} + // authViewer creates an authenticated viewer by validating the session key. func authViewer(ctx context.Context, sessionKey string, svc fleet.Service) (*viewer.Viewer, error) { session, err := svc.GetSessionByKey(ctx, sessionKey) diff --git a/server/service/handler.go b/server/service/handler.go index 30ed79634d..2118df9ec9 100644 --- a/server/service/handler.go +++ b/server/service/handler.go @@ -135,21 +135,21 @@ func MakeFleetServerEndpoints(svc fleet.Service, urlPrefix string, limitStore th throttled.RateQuota{MaxRate: throttled.PerMin(10), MaxBurst: 9})( makeLoginEndpoint(svc), ), - Logout: makeLogoutEndpoint(svc), + Logout: logged(makeLogoutEndpoint(svc)), ForgotPassword: limiter.Limit( throttled.RateQuota{MaxRate: throttled.PerHour(10), MaxBurst: 9})( - makeForgotPasswordEndpoint(svc), + logged(makeForgotPasswordEndpoint(svc)), ), - ResetPassword: makeResetPasswordEndpoint(svc), - CreateUserWithInvite: makeCreateUserFromInviteEndpoint(svc), - VerifyInvite: makeVerifyInviteEndpoint(svc), - InitiateSSO: makeInitiateSSOEndpoint(svc), - CallbackSSO: makeCallbackSSOEndpoint(svc, urlPrefix), - SSOSettings: makeSSOSettingsEndpoint(svc), + ResetPassword: logged(makeResetPasswordEndpoint(svc)), + CreateUserWithInvite: logged(makeCreateUserFromInviteEndpoint(svc)), + VerifyInvite: logged(makeVerifyInviteEndpoint(svc)), + InitiateSSO: logged(makeInitiateSSOEndpoint(svc)), + CallbackSSO: logged(makeCallbackSSOEndpoint(svc, urlPrefix)), + SSOSettings: logged(makeSSOSettingsEndpoint(svc)), // PerformRequiredPasswordReset needs only to authenticate the // logged in user - PerformRequiredPasswordReset: canPerformPasswordReset(makePerformRequiredPasswordResetEndpoint(svc)), + PerformRequiredPasswordReset: logged(canPerformPasswordReset(makePerformRequiredPasswordResetEndpoint(svc))), // Standard user authentication routes Me: authenticatedUser(svc, makeGetSessionUserEndpoint(svc)), @@ -242,7 +242,7 @@ func MakeFleetServerEndpoints(svc fleet.Service, urlPrefix string, limitStore th StatusLiveQuery: authenticatedUser(svc, makeStatusLiveQueryEndpoint(svc)), // Osquery endpoints - EnrollAgent: makeEnrollAgentEndpoint(svc), + EnrollAgent: logged(makeEnrollAgentEndpoint(svc)), // Authenticated osquery endpoints GetClientConfig: authenticatedHost(svc, makeGetClientConfigEndpoint(svc)), GetDistributedQueries: authenticatedHost(svc, makeGetDistributedQueriesEndpoint(svc)), @@ -252,7 +252,7 @@ func MakeFleetServerEndpoints(svc fleet.Service, urlPrefix string, limitStore th // For some reason osquery does not provide a node key with the block // data. Instead the carve session ID should be verified in the service // method. - CarveBlock: makeCarveBlockEndpoint(svc), + CarveBlock: logged(makeCarveBlockEndpoint(svc)), } } diff --git a/server/service/integration_test.go b/server/service/integration_test.go index d9f87d6a85..9c4ccbbd55 100644 --- a/server/service/integration_test.go +++ b/server/service/integration_test.go @@ -675,6 +675,39 @@ func TestVulnerableSoftware(t *testing.T) { assert.Contains(t, string(bodyBytes), expectedJSONSoft1) } +func TestOsqueryEndpointsLogErrors(t *testing.T) { + buf := new(bytes.Buffer) + logger := log.NewJSONLogger(buf) + logger = level.NewFilter(logger, level.AllowDebug()) + + ds := mysql.CreateMySQLDS(t) + defer ds.Close() + + _, server := RunServerForTestsWithDS(t, ds, TestServerOpts{Logger: logger}) + + _, err := ds.NewHost(&fleet.Host{ + DetailUpdatedAt: time.Now(), + LabelUpdatedAt: time.Now(), + SeenTime: time.Now(), + NodeKey: "1234", + UUID: "1", + Hostname: "foo.local", + PrimaryIP: "192.168.1.1", + PrimaryMac: "30-65-EC-6F-C4-58", + }) + require.NoError(t, err) + + requestBody := &nopCloser{bytes.NewBuffer([]byte(`{"node_key":"1234","log_type":"status","data":[}`))} + req, _ := http.NewRequest("POST", server.URL+"/api/v1/osquery/log", requestBody) + client := &http.Client{} + _, err = client.Do(req) + require.Nil(t, err) + + logString := buf.String() + assert.Equal(t, `{"err":"decoding JSON: invalid character '}' looking for beginning of value","level":"info","path":"/api/v1/osquery/log"} +`, logString) +} + func TestSubmitStatusLog(t *testing.T) { buf := new(bytes.Buffer) logger := log.NewJSONLogger(buf) @@ -710,3 +743,46 @@ func TestSubmitStatusLog(t *testing.T) { assert.Equal(t, 1, strings.Count(logString, "\"ip_addr\"")) assert.Equal(t, 1, strings.Count(logString, "x_for_ip_addr")) } + +func TestEnrollAgentLogsErrors(t *testing.T) { + buf := new(bytes.Buffer) + logger := log.NewJSONLogger(buf) + logger = level.NewFilter(logger, level.AllowDebug()) + + ds := mysql.CreateMySQLDS(t) + defer ds.Close() + + _, server := RunServerForTestsWithDS(t, ds, TestServerOpts{Logger: logger}) + + _, err := ds.NewHost(&fleet.Host{ + DetailUpdatedAt: time.Now(), + LabelUpdatedAt: time.Now(), + SeenTime: time.Now(), + NodeKey: "1234", + UUID: "1", + Hostname: "foo.local", + PrimaryIP: "192.168.1.1", + PrimaryMac: "30-65-EC-6F-C4-58", + }) + require.NoError(t, err) + + j, err := json.Marshal(&enrollAgentRequest{ + EnrollSecret: "1234", + HostIdentifier: "4321", + HostDetails: nil, + }) + require.NoError(t, err) + + requestBody := &nopCloser{bytes.NewBuffer(j)} + req, _ := http.NewRequest("POST", server.URL+"/api/v1/osquery/enroll", requestBody) + client := &http.Client{} + resp, err := client.Do(req) + require.NoError(t, err) + require.NoError(t, resp.Body.Close()) + + parts := strings.Split(strings.TrimSpace(buf.String()), "\n") + require.Len(t, parts, 1) + logData := make(map[string]json.RawMessage) + require.NoError(t, json.Unmarshal([]byte(parts[0]), &logData)) + assert.Equal(t, json.RawMessage(`["enroll failed: no matching secret found"]`), logData["err"]) +}