diff --git a/changes/29482-migrate-to-aws-sdk-go-v2 b/changes/29482-migrate-to-aws-sdk-go-v2 new file mode 100644 index 0000000000..20c35330db --- /dev/null +++ b/changes/29482-migrate-to-aws-sdk-go-v2 @@ -0,0 +1 @@ +* Migrated from `aws-sdk-go` v1 to `aws-sdk-go-v2`. diff --git a/cmd/fleet/main.go b/cmd/fleet/main.go index c60fc5a55a..8dc2ff2ebc 100644 --- a/cmd/fleet/main.go +++ b/cmd/fleet/main.go @@ -107,7 +107,7 @@ func applyDevFlags(cfg *config.FleetConfig) { cfg.S3.CarvesBucket = "carves-dev" cfg.S3.CarvesRegion = "minio" cfg.S3.CarvesPrefix = "dev-prefix" - cfg.S3.CarvesEndpointURL = "localhost:9000" + cfg.S3.CarvesEndpointURL = "http://localhost:9000" cfg.S3.CarvesAccessKeyID = "minio" cfg.S3.CarvesSecretAccessKey = "minio123!" cfg.S3.CarvesDisableSSL = true @@ -118,23 +118,12 @@ func applyDevFlags(cfg *config.FleetConfig) { cfg.S3.SoftwareInstallersBucket = "software-installers-dev" cfg.S3.SoftwareInstallersRegion = "minio" cfg.S3.SoftwareInstallersPrefix = "dev-prefix" - cfg.S3.SoftwareInstallersEndpointURL = "localhost:9000" + cfg.S3.SoftwareInstallersEndpointURL = "http://localhost:9000" cfg.S3.SoftwareInstallersAccessKeyID = "minio" cfg.S3.SoftwareInstallersSecretAccessKey = "minio123!" cfg.S3.SoftwareInstallersDisableSSL = true cfg.S3.SoftwareInstallersForceS3PathStyle = true } - - cfg.Packaging.S3 = config.S3Config{ - Bucket: "installers-dev", - Region: "minio", - Prefix: "dev-prefix", - EndpointURL: "localhost:9000", - AccessKeyID: "minio", - SecretAccessKey: "minio123!", - DisableSSL: true, - ForceS3PathStyle: true, - } } func initLogger(cfg config.FleetConfig) kitlog.Logger { diff --git a/cmd/fleet/serve.go b/cmd/fleet/serve.go index ea77f533d4..4e66a865af 100644 --- a/cmd/fleet/serve.go +++ b/cmd/fleet/serve.go @@ -138,7 +138,7 @@ the way that the Fleet server works. logger := initLogger(config) if dev { - createTestBucketForInstallers(&config, logger) + createTestBuckets(&config, logger) } // Init tracing @@ -254,41 +254,6 @@ the way that the Fleet server works. } } - if config.Packaging.GlobalEnrollSecret != "" { - secrets, err := ds.GetEnrollSecrets(cmd.Context(), nil) - if err != nil { - initFatal(err, "loading enroll secrets") - } - - var globalEnrollSecret string - for _, secret := range secrets { - if secret.TeamID == nil { - globalEnrollSecret = secret.Secret - break - } - } - - if globalEnrollSecret != "" { - if globalEnrollSecret != config.Packaging.GlobalEnrollSecret { - fmt.Printf("################################################################################\n" + - "# WARNING:\n" + - "# You have provided a global enroll secret config, but there's\n" + - "# already one set up for your application.\n" + - "#\n" + - "# This is generally an error and the provided value will be\n" + - "# ignored, if you really need to configure an enroll secret please\n" + - "# remove the global enroll secret from the database manually.\n" + - "################################################################################\n") - os.Exit(1) - } - } else { - if err := ds.ApplyEnrollSecrets(cmd.Context(), nil, - []*fleet.EnrollSecret{{Secret: config.Packaging.GlobalEnrollSecret}}); err != nil { - level.Debug(logger).Log("err", err, "msg", "failed to apply enroll secrets") - } - } - } - // Strip the Redis URI scheme if it's present. Scheme docs are at: https://www.iana.org/assignments/uri-schemes/uri-schemes.xhtml // This allows us to use Render's Redis service in render.yaml, including the free tier. // In the future, we could support the full Redis URI if needed (including username, password, database, etc.) @@ -1654,17 +1619,29 @@ func (n nopPusher) Push(context.Context, []string) (map[string]*push.Response, e return nil, nil } -func createTestBucketForInstallers(config *configpkg.FleetConfig, logger log.Logger) { - store, err := s3.NewSoftwareInstallerStore(config.S3) +func createTestBuckets(config *configpkg.FleetConfig, logger log.Logger) { + softwareInstallerStore, err := s3.NewSoftwareInstallerStore(config.S3) if err != nil { initFatal(err, "initializing S3 software installer store") } - if err := store.CreateTestBucket(config.S3.SoftwareInstallersBucket); err != nil { + if err := softwareInstallerStore.CreateTestBucket(context.Background(), config.S3.SoftwareInstallersBucket); err != nil { // Don't panic, allow devs to run Fleet without minio/S3 dependency. level.Info(logger).Log( "err", err, - "msg", "failed to create test bucket", + "msg", "failed to create test software installer bucket", "name", config.S3.SoftwareInstallersBucket, ) } + carveStore, err := s3.NewCarveStore(config.S3, nil) + if err != nil { + initFatal(err, "initializing S3 carve store") + } + if err := carveStore.CreateTestBucket(context.Background(), config.S3.CarvesBucket); err != nil { + // Don't panic, allow devs to run Fleet without minio/S3 dependency. + level.Info(logger).Log( + "err", err, + "msg", "failed to create test carve bucket", + "name", config.S3.CarvesBucket, + ) + } } diff --git a/docker-compose.yml b/docker-compose.yml index cc1279d944..a73ae97967 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -151,7 +151,7 @@ services: - "4566:4566" - "4571:4571" environment: - - SERVICES=firehose,kinesis + - SERVICES=firehose,kinesis,s3,iam,sts # s3 compatible object storage (file carving backend) minio: diff --git a/docs/Contributing/getting-started/testing-and-local-development.md b/docs/Contributing/getting-started/testing-and-local-development.md index 4edf6368a5..c60d2ca3e6 100644 --- a/docs/Contributing/getting-started/testing-and-local-development.md +++ b/docs/Contributing/getting-started/testing-and-local-development.md @@ -409,11 +409,13 @@ To add additional users, modify [tools/saml/users.php](https://github.com/fleetd -## Testing Kinesis Logging +## Testing Kinesis logging -Tip: Install [AwsLocal](https://github.com/localstack/awscli-local) to ease interaction with -[LocalStack](https://github.com/localstack/localstack). Alternatively, you can use the `aws` client -and use `--endpoint-url=http://localhost:4566` on all invocations. +Install the `aws` client: `brew install aws-cli` +Set the following alias to ease interaction with [LocalStack](https://github.com/localstack/localstack): +```sh +awslocal='AWS_ACCESS_KEY_ID=default AWS_SECRET_ACCESS_KEY=default AWS_DEFAULT_REGION=us-east-1 aws --endpoint-url=http://localhost:4566' +``` The following guide assumes you have server dependencies running: ```sh @@ -473,7 +475,6 @@ $ awslocal kinesis describe-stream --stream-name sample_status } ``` - Use the following configuration to run Fleet: ```sh FLEET_OSQUERY_RESULT_LOG_PLUGIN=kinesis @@ -545,26 +546,122 @@ echo eyJob3N0SWRlbnRpZmllciI6Ijg3OGE2ZWRmLTcxMzEtNGUyOC05NWEyLWQzNDQ5MDVjYWNhYiI {"hostIdentifier":"878a6edf-7131-4e28-95a2-d344905cacab","calendarTime":"Wed Mar 2 22:02:54 2022 UTC","unixTime":"1646258574","severity":"0","filename":"glog_logger.cpp","line":"49","message":"Could not get RPM header flag.","version":"4.9.0","decorations":{"host_uuid":"eb3946b2-0000-0000-b888-2591a1b666e9","hostname":"e0088d28a63f"}} ``` -## Testing pre-built installers +## Testing Firehose logging -Pre-built installers are kept in a blob storage like AWS S3. As part of your your local development there's a [MinIO](https://min.io/) instance running on http://localhost:9000. To test the pre-built installers functionality locally: - -1. Build the installers you want using `fleetctl package`. Be sure to include the `--insecure` flag - for local testing. -2. Use the [installerstore](https://github.com/fleetdm/fleet/tree/97b4d1f3fb30f7b25991412c0b40327f93cb118c/tools/installerstore) tool to upload them to your MinIO instance. -3. Configure your fleet server setting `FLEET_PACKAGING_GLOBAL_ENROLL_SECRET` to match your global enroll secret. -4. Set `FLEET_SERVER_SANDBOX_ENABLED=1`, as the endpoint to retrieve the installer is only available in the sandbox. +We will configure Fleet to send result and status logs to Firehose directly which will in turn stream them to S3 (`Fleet -> LocalStack Firehose -> LocalStack S3`). +Install the `aws` client: `brew install aws-cli` +Set the following alias to ease interaction with [LocalStack](https://github.com/localstack/localstack): ```sh -FLEET_SERVER_SANDBOX_ENABLED=1 FLEET_PACKAGING_GLOBAL_ENROLL_SECRET=xyz ./build/fleet serve --dev +awslocal='AWS_ACCESS_KEY_ID=test AWS_SECRET_ACCESS_KEY=test AWS_DEFAULT_REGION=us-east-1 aws --endpoint-url=http://localhost:4566' ``` -Be sure to replace the `FLEET_PACKAGING_GLOBAL_ENROLL_SECRET` value above with the global enroll -secret from the `fleetctl package` command used to build the installers. +We need to create a S3 bucket in LocalStack and make it "publicly" available (so that we can inspect it in the browser) +```sh +awslocal s3 mb s3://s3-firehose --region us-east-1 +awslocal s3api put-bucket-acl --bucket s3-firehose --acl public-read +``` +Check `http://localhost:4566/s3-firehose` in your browser. -MinIO also offers a web interface at http://localhost:9001. Credentials are `minio` / `minio123!`. When starting the -Fleet server up with `--dev` the server will look for installers in the `software-installers-dev` MinIO bucket. You can -create this bucket via the MinIO web UI (it is *not* created by default when setting up the docker-compose environment). +Create the following `iam_policy.json` file and apply it to create a "super-role": +```json +{ + "Version": "2012-10-17", + "Statement": [ + { + "Sid": "Stmt1572416334166", + "Action": "*", + "Effect": "Allow", + "Resource": "*" + } + ] +} +``` +```sh +awslocal iam create-role --role-name super-role --assume-role-policy-document file://$(pwd)/iam_policy.json +``` + +After applying it, grab the "Arn" in the output (e.g. `"arn:aws:iam::000000000000:role/super-role"`) + +Create the following `firehose_skeleton_result.json` file to create the delivery stream for "result" logs: +```json +{ + "DeliveryStreamName": "s3-stream-result", + "DeliveryStreamType": "DirectPut", + "S3DestinationConfiguration": { + "RoleARN": "arn:aws:iam::000000000000:role/super-role", + "BucketARN": "arn:aws:s3:::s3-firehose", + "Prefix": "result", + "ErrorOutputPrefix": "result-error", + "BufferingHints": { + "SizeInMBs": 1, + "IntervalInSeconds": 60 + }, + "CompressionFormat": "UNCOMPRESSED", + "CloudWatchLoggingOptions": { + "Enabled": false, + "LogGroupName": "", + "LogStreamName": "" + } + }, + "Tags": [ + { + "Key": "tagKey", + "Value": "tagValue" + } + ] +} +``` +```sh +awslocal firehose create-delivery-stream --cli-input-json file://$(pwd)/firehose_skeleton_result.json +``` + +Similarly, create a `firehose_skeleton_status.json` file to create the delivery stream for "status" logs: +```json +{ + "DeliveryStreamName": "s3-stream-status", + "DeliveryStreamType": "DirectPut", + "S3DestinationConfiguration": { + "RoleARN": "arn:aws:iam::000000000000:role/super-role", + "BucketARN": "arn:aws:s3:::s3-firehose", + "Prefix": "status", + "ErrorOutputPrefix": "status-error", + "BufferingHints": { + "SizeInMBs": 1, + "IntervalInSeconds": 60 + }, + "CompressionFormat": "UNCOMPRESSED", + "CloudWatchLoggingOptions": { + "Enabled": false, + "LogGroupName": "", + "LogStreamName": "" + } + }, + "Tags": [ + { + "Key": "tagKey", + "Value": "tagValue" + } + ] +} +``` + +After applying such configuration, "result" logs will be stored under the `results/` prefix and "status" logs will be stored under `status/` prefix (both on the `s3-firehose` bucket). + +Finally, here's the Fleet configuration: +```sh +FLEET_OSQUERY_RESULT_LOG_PLUGIN=firehose +FLEET_OSQUERY_STATUS_LOG_PLUGIN=firehose +FLEET_FIREHOSE_REGION=us-east-1 +FLEET_FIREHOSE_ENDPOINT_URL=http://localhost:4566 +FLEET_FIREHOSE_ACCESS_KEY_ID=default +FLEET_FIREHOSE_SECRET_ACCESS_KEY=default +FLEET_FIREHOSE_STS_ASSUME_ROLE_ARN=arn:aws:iam::000000000000:role/super-role +FLEET_FIREHOSE_RESULT_STREAM=s3-stream-result +FLEET_FIREHOSE_STATUS_STREAM=s3-stream-status +``` + +You can inspect logs by visiting `http://localhost:4566/s3-firehose` on your browser. ## Telemetry @@ -943,7 +1040,7 @@ open /opt/orbit/bin/nudge/macos/stable/Nudge.app --args -json-url file:///opt/or ### Bootstrap package -A bootstrap package is a `pkg` file that gets automatically installed on hosts when they enroll via DEP. +A bootstrap package is a `pkg` file that gets automatically installed on hosts when they enroll via ABM/DEP. The `pkg` file needs to be a signed "distribution package", you can find a dummy file that meets all the requirements [in Drive](https://drive.google.com/file/d/1adwAOTD5G6D4WzWvJeMId6mDhyeFy-lm/view). We have instructions in [the docs](https://fleetdm.com/docs/using-fleet/mdm-macos-setup-experience#bootstrap-package) to upload a new bootstrap package to your Fleet instance. diff --git a/docs/Contributing/product-groups/orchestration/file-carving.md b/docs/Contributing/product-groups/orchestration/file-carving.md index b283f11f4c..d35aa773df 100644 --- a/docs/Contributing/product-groups/orchestration/file-carving.md +++ b/docs/Contributing/product-groups/orchestration/file-carving.md @@ -83,14 +83,11 @@ Configure the following: - `FLEET_S3_ENDPOINT_URL=minio_host:port` - `FLEET_S3_BUCKET=minio_bucket_name` - `FLEET_S3_SECRET_ACCESS_KEY=your_secret_access_key` -- `FLEET_S3_ACCESS_KEY_ID=acces_key_id` +- `FLEET_S3_ACCESS_KEY_ID=access_key_id` - `FLEET_S3_FORCE_S3_PATH_STYLE=true` - `FLEET_S3_REGION=minio` or any non-empty string otherwise Fleet will attempt to derive the region. -If you're testing file carving locally with the docker-compose environment, the `--dev` flag on Fleet server will -automatically point carves to the local MinIO container and write to the `carves-dev` bucket without needing to set -additional configuration. Note that this bucket is *not* created automatically when bringing MinIO up; you'll need to -log in via `http://localhost:9001` with credentials `minio` / `minio123!` to create the bucket. +If you're testing file carving locally with the docker-compose environment, the `--dev` flag on Fleet server will automatically point carves to the local MinIO container and write to the `carves-dev` bucket (created automatically) without needing to set additional configuration. ### Troubleshooting diff --git a/go.mod b/go.mod index d13eeccf51..29c7d8d3a3 100644 --- a/go.mod +++ b/go.mod @@ -17,8 +17,17 @@ require ( github.com/andygrunwald/go-jira v1.16.0 github.com/antchfx/xmlquery v1.3.14 github.com/apex/log v1.9.0 - github.com/aws/aws-sdk-go v1.44.288 + github.com/aws/aws-sdk-go-v2 v1.36.5 + github.com/aws/aws-sdk-go-v2/config v1.29.17 + github.com/aws/aws-sdk-go-v2/credentials v1.17.70 github.com/aws/aws-sdk-go-v2/feature/cloudfront/sign v1.8.3 + github.com/aws/aws-sdk-go-v2/feature/s3/manager v1.17.81 + github.com/aws/aws-sdk-go-v2/service/firehose v1.37.7 + github.com/aws/aws-sdk-go-v2/service/kinesis v1.35.3 + github.com/aws/aws-sdk-go-v2/service/lambda v1.72.0 + github.com/aws/aws-sdk-go-v2/service/s3 v1.81.0 + github.com/aws/aws-sdk-go-v2/service/ses v1.30.4 + github.com/aws/smithy-go v1.22.4 github.com/beevik/etree v1.3.0 github.com/beevik/ntp v0.3.0 github.com/blakesmith/ar v0.0.0-20190502131153-809d4375e1fb @@ -176,6 +185,19 @@ require ( github.com/apache/thrift v0.18.1 // indirect github.com/armon/circbuf v0.0.0-20190214190532-5111143e8da2 // indirect github.com/armon/go-radix v1.0.0 // indirect + github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.6.11 // indirect + github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.16.32 // indirect + github.com/aws/aws-sdk-go-v2/internal/configsources v1.3.36 // indirect + github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.6.36 // indirect + github.com/aws/aws-sdk-go-v2/internal/ini v1.8.3 // indirect + github.com/aws/aws-sdk-go-v2/internal/v4a v1.3.36 // indirect + github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.12.4 // indirect + github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.7.4 // indirect + github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.12.17 // indirect + github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.18.17 // indirect + github.com/aws/aws-sdk-go-v2/service/sso v1.25.5 // indirect + github.com/aws/aws-sdk-go-v2/service/ssooidc v1.30.3 // indirect + github.com/aws/aws-sdk-go-v2/service/sts v1.34.0 // indirect github.com/beorn7/perks v1.0.1 // indirect github.com/c-bata/go-prompt v0.2.3 // indirect github.com/cavaliercoder/go-cpio v0.0.0-20180626203310-925f9528c45e // indirect @@ -238,7 +260,6 @@ require ( github.com/imdario/mergo v0.3.12 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/jbenet/go-context v0.0.0-20150711004518-d14ea06fba99 // indirect - github.com/jmespath/go-jmespath v0.4.0 // indirect github.com/joeshaw/multierror v0.0.0-20140124173710-69b34d4ec901 // indirect github.com/jonboulle/clockwork v0.2.2 // indirect github.com/kevinburke/go-bindata v3.24.0+incompatible // indirect diff --git a/go.sum b/go.sum index c775514e24..182fd8775f 100644 --- a/go.sum +++ b/go.sum @@ -131,14 +131,54 @@ github.com/armon/go-radix v1.0.0/go.mod h1:ufUuZ+zHj4x4TnLV4JWEpy2hxWSpsRywHrMgI github.com/armon/go-socks5 v0.0.0-20160902184237-e75332964ef5 h1:0CwZNZbxp69SHPdPJAN/hZIm0C4OItdklCFmMRWYpio= github.com/armon/go-socks5 v0.0.0-20160902184237-e75332964ef5/go.mod h1:wHh0iHkYZB8zMSxRWpUBQtwG5a7fFgvEO+odwuTv2gs= github.com/aws/aws-sdk-go v1.20.6/go.mod h1:KmX6BPdI08NWTb3/sm4ZGu5ShLoqVDhKgpiN924inxo= -github.com/aws/aws-sdk-go v1.44.288 h1:Ln7fIao/nl0ACtelgR1I4AiEw/GLNkKcXfCaHupUW5Q= -github.com/aws/aws-sdk-go v1.44.288/go.mod h1:aVsgQcEevwlmQ7qHE9I3h+dtQgpqhFB+i8Phjh7fkwI= -github.com/aws/aws-sdk-go-v2 v1.32.7 h1:ky5o35oENWi0JYWUZkB7WYvVPP+bcRF5/Iq7JWSb5Rw= -github.com/aws/aws-sdk-go-v2 v1.32.7/go.mod h1:P5WJBrYqqbWVaOxgH0X/FYYD47/nooaPOZPlQdmiN2U= +github.com/aws/aws-sdk-go-v2 v1.36.5 h1:0OF9RiEMEdDdZEMqF9MRjevyxAQcf6gY+E7vwBILFj0= +github.com/aws/aws-sdk-go-v2 v1.36.5/go.mod h1:EYrzvCCN9CMUTa5+6lf6MM4tq3Zjp8UhSGR/cBsjai0= +github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.6.11 h1:12SpdwU8Djs+YGklkinSSlcrPyj3H4VifVsKf78KbwA= +github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.6.11/go.mod h1:dd+Lkp6YmMryke+qxW/VnKyhMBDTYP41Q2Bb+6gNZgY= +github.com/aws/aws-sdk-go-v2/config v1.29.17 h1:jSuiQ5jEe4SAMH6lLRMY9OVC+TqJLP5655pBGjmnjr0= +github.com/aws/aws-sdk-go-v2/config v1.29.17/go.mod h1:9P4wwACpbeXs9Pm9w1QTh6BwWwJjwYvJ1iCt5QbCXh8= +github.com/aws/aws-sdk-go-v2/credentials v1.17.70 h1:ONnH5CM16RTXRkS8Z1qg7/s2eDOhHhaXVd72mmyv4/0= +github.com/aws/aws-sdk-go-v2/credentials v1.17.70/go.mod h1:M+lWhhmomVGgtuPOhO85u4pEa3SmssPTdcYpP/5J/xc= github.com/aws/aws-sdk-go-v2/feature/cloudfront/sign v1.8.3 h1:/d7ZHq/2m+1Uzw4mnizCZbTAWB/dJ3CPy0N1qUpUpI0= github.com/aws/aws-sdk-go-v2/feature/cloudfront/sign v1.8.3/go.mod h1:xWMYk6dLhV33jy2YrbOsv2l3fZTDMWE1yIIbvnD13gU= -github.com/aws/smithy-go v1.22.1 h1:/HPHZQ0g7f4eUeK6HKglFz8uwVfZKgoI25rb/J+dnro= -github.com/aws/smithy-go v1.22.1/go.mod h1:irrKGvNn1InZwb2d7fkIRNucdfwR8R+Ts3wxYa/cJHg= +github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.16.32 h1:KAXP9JSHO1vKGCr5f4O6WmlVKLFFXgWYAGoJosorxzU= +github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.16.32/go.mod h1:h4Sg6FQdexC1yYG9RDnOvLbW1a/P986++/Y/a+GyEM8= +github.com/aws/aws-sdk-go-v2/feature/s3/manager v1.17.81 h1:E5ff1vZlAudg24j5lF6F6/gBpln2LjWxGdQDBSLfVe4= +github.com/aws/aws-sdk-go-v2/feature/s3/manager v1.17.81/go.mod h1:hHBLCuhHI4Aokvs5vdVoCDBzmFy86yxs5J7LEPQwQEM= +github.com/aws/aws-sdk-go-v2/internal/configsources v1.3.36 h1:SsytQyTMHMDPspp+spo7XwXTP44aJZZAC7fBV2C5+5s= +github.com/aws/aws-sdk-go-v2/internal/configsources v1.3.36/go.mod h1:Q1lnJArKRXkenyog6+Y+zr7WDpk4e6XlR6gs20bbeNo= +github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.6.36 h1:i2vNHQiXUvKhs3quBR6aqlgJaiaexz/aNvdCktW/kAM= +github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.6.36/go.mod h1:UdyGa7Q91id/sdyHPwth+043HhmP6yP9MBHgbZM0xo8= +github.com/aws/aws-sdk-go-v2/internal/ini v1.8.3 h1:bIqFDwgGXXN1Kpp99pDOdKMTTb5d2KyU5X/BZxjOkRo= +github.com/aws/aws-sdk-go-v2/internal/ini v1.8.3/go.mod h1:H5O/EsxDWyU+LP/V8i5sm8cxoZgc2fdNR9bxlOFrQTo= +github.com/aws/aws-sdk-go-v2/internal/v4a v1.3.36 h1:GMYy2EOWfzdP3wfVAGXBNKY5vK4K8vMET4sYOYltmqs= +github.com/aws/aws-sdk-go-v2/internal/v4a v1.3.36/go.mod h1:gDhdAV6wL3PmPqBhiPbnlS447GoWs8HTTOYef9/9Inw= +github.com/aws/aws-sdk-go-v2/service/firehose v1.37.7 h1:rDNxf0CQboBMqzm6WmhGL58pYpKMjU6Qs3/BfY3Em4Y= +github.com/aws/aws-sdk-go-v2/service/firehose v1.37.7/go.mod h1:E1yDRkUMwlVGmDYcu5UJuwfznGNuVW29sjr2xxM2Y0w= +github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.12.4 h1:CXV68E2dNqhuynZJPB80bhPQwAKqBWVer887figW6Jc= +github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.12.4/go.mod h1:/xFi9KtvBXP97ppCz1TAEvU1Uf66qvid89rbem3wCzQ= +github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.7.4 h1:nAP2GYbfh8dd2zGZqFRSMlq+/F6cMPBUuCsGAMkN074= +github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.7.4/go.mod h1:LT10DsiGjLWh4GbjInf9LQejkYEhBgBCjLG5+lvk4EE= +github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.12.17 h1:t0E6FzREdtCsiLIoLCWsYliNsRBgyGD/MCK571qk4MI= +github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.12.17/go.mod h1:ygpklyoaypuyDvOM5ujWGrYWpAK3h7ugnmKCU/76Ys4= +github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.18.17 h1:qcLWgdhq45sDM9na4cvXax9dyLitn8EYBRl8Ak4XtG4= +github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.18.17/go.mod h1:M+jkjBFZ2J6DJrjMv2+vkBbuht6kxJYtJiwoVgX4p4U= +github.com/aws/aws-sdk-go-v2/service/kinesis v1.35.3 h1:aAi9YBNpYMEX52Z9qy1YP2t3RhDqMcP67Ep/C4q5RiQ= +github.com/aws/aws-sdk-go-v2/service/kinesis v1.35.3/go.mod h1:DH0TzTbBG82HKNpBQlplRNSS4bGz0dsbJvxdK9f6rUY= +github.com/aws/aws-sdk-go-v2/service/lambda v1.72.0 h1:2LerDz2Lz22IDfdpR/RpSZIFoBoAh1tdHUaiUzG2z0k= +github.com/aws/aws-sdk-go-v2/service/lambda v1.72.0/go.mod h1:vahA7MiX/fQE9J5o1PKbgn8KoXz7ogSFLAQQLdLUvM8= +github.com/aws/aws-sdk-go-v2/service/s3 v1.81.0 h1:1GmCadhKR3J2sMVKs2bAYq9VnwYeCqfRyZzD4RASGlA= +github.com/aws/aws-sdk-go-v2/service/s3 v1.81.0/go.mod h1:kUklwasNoCn5YpyAqC/97r6dzTA1SRKJfKq16SXeoDU= +github.com/aws/aws-sdk-go-v2/service/ses v1.30.4 h1:VT+yYtHKQiDJrNAsvoO2ExMUN3KxWsFRt+S5j1MdFGk= +github.com/aws/aws-sdk-go-v2/service/ses v1.30.4/go.mod h1:Zftob00wu8O9xWSN1pdczm1U+E6yXk9znf+4lkt+3aQ= +github.com/aws/aws-sdk-go-v2/service/sso v1.25.5 h1:AIRJ3lfb2w/1/8wOOSqYb9fUKGwQbtysJ2H1MofRUPg= +github.com/aws/aws-sdk-go-v2/service/sso v1.25.5/go.mod h1:b7SiVprpU+iGazDUqvRSLf5XmCdn+JtT1on7uNL6Ipc= +github.com/aws/aws-sdk-go-v2/service/ssooidc v1.30.3 h1:BpOxT3yhLwSJ77qIY3DoHAQjZsc4HEGfMCE4NGy3uFg= +github.com/aws/aws-sdk-go-v2/service/ssooidc v1.30.3/go.mod h1:vq/GQR1gOFLquZMSrxUK/cpvKCNVYibNyJ1m7JrU88E= +github.com/aws/aws-sdk-go-v2/service/sts v1.34.0 h1:NFOJ/NXEGV4Rq//71Hs1jC/NvPs1ezajK+yQmkwnPV0= +github.com/aws/aws-sdk-go-v2/service/sts v1.34.0/go.mod h1:7ph2tGpfQvwzgistp2+zga9f+bCjlQJPkPUmMgDSD7w= +github.com/aws/smithy-go v1.22.4 h1:uqXzVZNuNexwc/xrh6Tb56u89WDlJY6HS+KC0S4QSjw= +github.com/aws/smithy-go v1.22.4/go.mod h1:t1ufH5HMublsJYulve2RKmHDC15xu1f26kHCp/HgceI= github.com/aybabtme/rgbterm v0.0.0-20170906152045-cc83f3b3ce59/go.mod h1:q/89r3U2H7sSsE2t6Kca0lfwTK8JdoNGS/yzM/4iH5I= github.com/beevik/etree v1.1.0/go.mod h1:r8Aw8JqVegEf0w2fDnATrX9VpkMcyFeM0FhwO62wh+A= github.com/beevik/etree v1.3.0 h1:hQTc+pylzIKDb23yYprodCWWTt+ojFfUZyzU09a/hmU= @@ -544,10 +584,6 @@ github.com/jbenet/go-context v0.0.0-20150711004518-d14ea06fba99 h1:BQSFePA1RWJOl github.com/jbenet/go-context v0.0.0-20150711004518-d14ea06fba99/go.mod h1:1lJo3i6rXxKeerYnT8Nvf0QmHCRC1n8sfWVwXF2Frvo= github.com/jessevdk/go-flags v1.4.0/go.mod h1:4FA24M0QyGHXBuZZK/XkWh8h0e1EYbRYJSGM75WSRxI= github.com/jmespath/go-jmespath v0.0.0-20180206201540-c2b33e8439af/go.mod h1:Nht3zPeWKUH0NzdCt2Blrr5ys8VGpn0CEB0cQHVjt7k= -github.com/jmespath/go-jmespath v0.4.0 h1:BEgLn5cpjn8UN1mAw4NjwDrS35OdebyEtFe+9YPoQUg= -github.com/jmespath/go-jmespath v0.4.0/go.mod h1:T8mJZnbsbmF+m6zOOFylbeCJqk5+pHWvzYPziyZiYoo= -github.com/jmespath/go-jmespath/internal/testify v1.5.1 h1:shLQSRRSCCPj3f2gpwzGwWFoC7ycTf1rcQZHOlsJ6N8= -github.com/jmespath/go-jmespath/internal/testify v1.5.1/go.mod h1:L3OGu8Wl2/fWfCI6z80xFu9LTZmf1ZRjMHUOPmWr69U= github.com/jmoiron/sqlx v0.0.0-20180406164412-2aeb6a910c2b/go.mod h1:IiEW3SEiiErVyFdH8NTuWjSifiEQKUoyK3LNqr2kCHU= github.com/jmoiron/sqlx v1.3.5 h1:vFFPA71p1o5gAeqtEAwLU4dnX2napprKtHr7PYIcN3g= github.com/jmoiron/sqlx v1.3.5/go.mod h1:nRVWtLre0KfCLJvgxzCsLVMogSvQ1zNJtpYr2Ccp0mQ= @@ -1090,7 +1126,6 @@ golang.org/x/net v0.0.0-20210405180319-a5a99cb37ef4/go.mod h1:p54w0d4576C0XHj96b golang.org/x/net v0.0.0-20210614182718-04defd469f4e/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y= golang.org/x/net v0.0.0-20211112202133-69e39bad7dc2/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y= golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c= -golang.org/x/net v0.1.0/go.mod h1:Cx3nUiGt4eDBEyega/BKRp+/AlGL8hYe7U9odMt2Cco= golang.org/x/net v0.5.0/go.mod h1:DivGGAXEgPSlEBzxGzZI+ZLohi+xUj054jfeKui00ws= golang.org/x/net v0.38.0 h1:vRMAPTMaeGqVhG5QyLJHqNDwecKTomGeqbnfZyKlBI8= golang.org/x/net v0.38.0/go.mod h1:ivrbrMbzFq5J41QOQh0siUuly180yBYtLp+CKbEaFx8= @@ -1161,7 +1196,6 @@ golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBc golang.org/x/sys v0.0.0-20220715151400-c0bba94af5f8/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220722155257-8c9f86f7a55f/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220811171246-fbc7d0a398ab/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.1.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.4.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.8.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= @@ -1173,7 +1207,6 @@ golang.org/x/sys v0.32.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k= golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= golang.org/x/term v0.0.0-20210220032956-6a3ed077a48d/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8= -golang.org/x/term v0.1.0/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8= golang.org/x/term v0.4.0/go.mod h1:9P2UbLfCdcvo3p/nzKvsmas4TnlujnuoV9hGgYzW1lQ= golang.org/x/term v0.30.0 h1:PQ39fJZ+mfadBm0y5WlL4vlM7Sx1Hgf13sMIY2+QS9Y= golang.org/x/term v0.30.0/go.mod h1:NYYFdzHoI5wRh/h5tDMdMqCqPJZEuNqVR5xJLd/n67g= @@ -1184,7 +1217,6 @@ golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.4/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ= -golang.org/x/text v0.4.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8= golang.org/x/text v0.6.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8= golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8= golang.org/x/text v0.23.0 h1:D71I7dUrlY+VX0gQShAThNGHFxZ13dGLBHQLVl1mJlY= diff --git a/security/code/.trivyignore b/security/code/.trivyignore index 229b985973..676e86fc23 100644 --- a/security/code/.trivyignore +++ b/security/code/.trivyignore @@ -1,10 +1,3 @@ -# These AWS SDK CVEs do not impact Fleet as we do not use S3 client-side crypto features - -CVE-2020-8911 -CVE-2020-8912 -GHSA-7f33-f4f5-xwgw -GHSA-f5pg-7wfw-84q9 - # Vulnerable code in trim is not used in Fleet CVE-2020-7753 diff --git a/server/config/config.go b/server/config/config.go index 2d8d86bd8e..cf83660336 100644 --- a/server/config/config.go +++ b/server/config/config.go @@ -578,6 +578,8 @@ type HTTPBasicAuthConfig struct { } // PackagingConfig holds configuration to build and retrieve Fleet packages +// +// Deprecated: "packaging" fields were used for "Fleet Sandbox" which doesn't exist anymore. type PackagingConfig struct { // GlobalEnrollSecret is the enroll secret that will be used to enroll // hosts in the global scope @@ -617,11 +619,13 @@ type FleetConfig struct { Sentry SentryConfig GeoIP GeoIPConfig Prometheus PrometheusConfig - Packaging PackagingConfig MDM MDMConfig Calendar CalendarConfig Partnerships PartnershipsConfig MicrosoftCompliancePartner MicrosoftCompliancePartnerConfig `yaml:"microsoft_compliance_partner"` + + // Deprecated: "packaging" fields were used for "Fleet Sandbox" which doesn't exist anymore. + Packaging PackagingConfig } type PartnershipsConfig struct { @@ -1394,6 +1398,8 @@ func (man Manager) addConfigs() { man.addConfigBool("prometheus.basic_auth.disable", false, "Disable HTTP Basic Auth for Prometheus") // Packaging config + // + // DEPRECATED: "packaging" fields were used for "Fleet Sandbox" which doesn't exist anymore. man.addConfigString("packaging.global_enroll_secret", "", "Enroll secret to be used for the global domain (instead of randomly generating one)") man.addConfigString("packaging.s3.bucket", "", "Bucket where to retrieve installers") man.addConfigString("packaging.s3.prefix", "", "Prefix under which installers are stored") @@ -1684,21 +1690,6 @@ func (man Manager) LoadConfig() FleetConfig { Disable: man.getConfigBool("prometheus.basic_auth.disable"), }, }, - Packaging: PackagingConfig{ - GlobalEnrollSecret: man.getConfigString("packaging.global_enroll_secret"), - S3: S3Config{ - Bucket: man.getConfigString("packaging.s3.bucket"), - Prefix: man.getConfigString("packaging.s3.prefix"), - Region: man.getConfigString("packaging.s3.region"), - EndpointURL: man.getConfigString("packaging.s3.endpoint_url"), - AccessKeyID: man.getConfigString("packaging.s3.access_key_id"), - SecretAccessKey: man.getConfigString("packaging.s3.secret_access_key"), - StsAssumeRoleArn: man.getConfigString("packaging.s3.sts_assume_role_arn"), - StsExternalID: man.getConfigString("packaging.s3.sts_external_id"), - DisableSSL: man.getConfigBool("packaging.s3.disable_ssl"), - ForceS3PathStyle: man.getConfigBool("packaging.s3.force_s3_path_style"), - }, - }, MDM: MDMConfig{ AppleAPNsCert: man.getConfigString("mdm.apple_apns_cert"), AppleAPNsCertBytes: man.getConfigString("mdm.apple_apns_cert_bytes"), diff --git a/server/config/config_test.go b/server/config/config_test.go index 3bcd99c527..af2386f0a5 100644 --- a/server/config/config_test.go +++ b/server/config/config_test.go @@ -70,6 +70,9 @@ func TestConfigRoundtrip(t *testing.T) { // These are deprecated field names in the S3 config. Set them to zero value, which leads to the new fields being populated instead. case "Bucket", "Prefix", "Region", "EndpointURL", "AccessKeyID", "SecretAccessKey", "StsAssumeRoleArn", "StsExternalID": key_v.SetString("") + // This is a deprecated config for "Fleet Sandbox" that doesn't exist anymore. + case "GlobalEnrollSecret": + key_v.SetString("") default: key_v.SetString(v.Elem().Type().Field(conf_index).Name + "_" + conf_v.Type().Field(key_index).Name) } diff --git a/server/datastore/s3/bootstrap_package.go b/server/datastore/s3/bootstrap_package.go index 28a41f51ef..8e893bde83 100644 --- a/server/datastore/s3/bootstrap_package.go +++ b/server/datastore/s3/bootstrap_package.go @@ -11,7 +11,7 @@ type BootstrapPackageStore struct { // NewBootstrapPackageStore creates a new instance with the given S3 config. func NewBootstrapPackageStore(config config.S3Config) (*BootstrapPackageStore, error) { // bootstrap packages use the same S3 config as software installers - s3store, err := newS3store(config.SoftwareInstallersToInternalCfg()) + s3store, err := newS3Store(config.SoftwareInstallersToInternalCfg()) if err != nil { return nil, err } diff --git a/server/datastore/s3/bootstrap_package_test.go b/server/datastore/s3/bootstrap_package_test.go index 11c4b1fe61..10c4824cbb 100644 --- a/server/datastore/s3/bootstrap_package_test.go +++ b/server/datastore/s3/bootstrap_package_test.go @@ -12,7 +12,7 @@ import ( "testing" "time" - "github.com/aws/aws-sdk-go/service/s3" + "github.com/aws/aws-sdk-go-v2/service/s3" "github.com/fleetdm/fleet/v4/server/fleet" "github.com/google/uuid" "github.com/stretchr/testify/require" @@ -91,7 +91,7 @@ func TestBootstrapPackageCleanup(t *testing.T) { assertExisting := func(want []string) { prefix := path.Join(store.prefix, bootstrapPackagePrefix) - page, err := store.s3client.ListObjectsV2(&s3.ListObjectsV2Input{ + page, err := store.s3Client.ListObjectsV2(ctx, &s3.ListObjectsV2Input{ Bucket: &store.bucket, Prefix: &prefix, }) diff --git a/server/datastore/s3/carves.go b/server/datastore/s3/carves.go index 7aae963301..17b8e33e4d 100644 --- a/server/datastore/s3/carves.go +++ b/server/datastore/s3/carves.go @@ -6,11 +6,12 @@ import ( "errors" "fmt" "io" + "strconv" "strings" "time" - "github.com/aws/aws-sdk-go/aws/awserr" - "github.com/aws/aws-sdk-go/service/s3" + "github.com/aws/aws-sdk-go-v2/service/s3" + "github.com/aws/aws-sdk-go-v2/service/s3/types" "github.com/fleetdm/fleet/v4/server/config" "github.com/fleetdm/fleet/v4/server/contexts/ctxerr" "github.com/fleetdm/fleet/v4/server/fleet" @@ -35,7 +36,7 @@ type CarveStore struct { // NewCarveStore creates a new store with the given config func NewCarveStore(config config.S3Config, metadatadb fleet.CarveStore) (*CarveStore, error) { - s3store, err := newS3store(config.CarvesToInternalCfg()) + s3store, err := newS3Store(config.CarvesToInternalCfg()) if err != nil { return nil, err } @@ -53,7 +54,7 @@ func (c *CarveStore) generateS3Key(metadata *fleet.CarveMetadata) string { // NewCarve initializes a new file carving session func (c *CarveStore) NewCarve(ctx context.Context, metadata *fleet.CarveMetadata) (*fleet.CarveMetadata, error) { objectKey := c.generateS3Key(metadata) - res, err := c.s3client.CreateMultipartUpload(&s3.CreateMultipartUploadInput{ + res, err := c.s3Client.CreateMultipartUpload(ctx, &s3.CreateMultipartUploadInput{ Bucket: &c.bucket, Key: &objectKey, }) @@ -84,7 +85,7 @@ func (c *CarveStore) UpdateCarve(ctx context.Context, metadata *fleet.CarveMetad // listS3Carves lists all keys up to a given one or if the passed max number // of keys has been reached; keys are returned in a set-like map -func (c *CarveStore) listS3Carves(lastPrefix string, maxKeys int) (map[string]bool, error) { +func (c *CarveStore) listS3Carves(ctx context.Context, lastPrefix string, maxKeys int) (map[string]bool, error) { var err error var continuationToken string result := make(map[string]bool) @@ -95,7 +96,7 @@ func (c *CarveStore) listS3Carves(lastPrefix string, maxKeys int) (map[string]bo lastPrefix = c.prefix + lastPrefix } for { - carveFilesPage, err := c.s3client.ListObjectsV2(&s3.ListObjectsV2Input{ + carveFilesPage, err := c.s3Client.ListObjectsV2(ctx, &s3.ListObjectsV2Input{ Bucket: &c.bucket, Prefix: &c.prefix, ContinuationToken: &continuationToken, @@ -134,7 +135,7 @@ func (c *CarveStore) CleanupCarves(ctx context.Context, now time.Time) (int, err // List carves in S3 up to a hour+1 prefix lastCarveNextHour := nonExpiredCarves[len(nonExpiredCarves)-1].CreatedAt.Add(time.Hour) lastCarvePrefix := c.prefix + lastCarveNextHour.Format(timePrefixFormat) - carveKeys, err := c.listS3Carves(lastCarvePrefix, 2*cleanupSize) + carveKeys, err := c.listS3Carves(ctx, lastCarvePrefix, 2*cleanupSize) if err != nil { return 0, ctxerr.Wrap(ctx, err, "s3 carve cleanup") } @@ -172,21 +173,23 @@ func (c *CarveStore) ListCarves(ctx context.Context, opt fleet.CarveListOptions) // listCompletedParts returns a list of the parts in a multipart updaload given a key and uploadID // results are wrapped into the s3.CompletedPart struct -func (c *CarveStore) listCompletedParts(objectKey, uploadID string) ([]*s3.CompletedPart, error) { - var res []*s3.CompletedPart - var partMarker int64 +func (c *CarveStore) listCompletedParts(ctx context.Context, objectKey, uploadID string) ([]types.CompletedPart, error) { + var res []types.CompletedPart + var partMarker int32 + for { - parts, err := c.s3client.ListParts(&s3.ListPartsInput{ + partNumberMarker := fmt.Sprint(partMarker) + parts, err := c.s3Client.ListParts(ctx, &s3.ListPartsInput{ Bucket: &c.bucket, Key: &objectKey, UploadId: &uploadID, - PartNumberMarker: &partMarker, + PartNumberMarker: &partNumberMarker, }) if err != nil { - return res, err + return nil, err } for _, p := range parts.Parts { - res = append(res, &s3.CompletedPart{ + res = append(res, types.CompletedPart{ ETag: p.ETag, PartNumber: p.PartNumber, }) @@ -194,16 +197,26 @@ func (c *CarveStore) listCompletedParts(objectKey, uploadID string) ([]*s3.Compl if !*parts.IsTruncated { break } - partMarker = *parts.NextPartNumberMarker + pm, err := strconv.ParseInt(*parts.NextPartNumberMarker, 10, 32) + if err != nil { + return nil, fmt.Errorf("failed to parse next part number marker: %q: %w", *parts.NextPartNumberMarker, err) + } + partMarker = int32(pm) } return res, nil } +const maxPartNumber = 10_000 + // NewBlock uploads a new block for a specific carve func (c *CarveStore) NewBlock(ctx context.Context, metadata *fleet.CarveMetadata, blockID int64, data []byte) error { + if blockID < 0 || blockID >= maxPartNumber { + return ctxerr.Errorf(ctx, "invalid blockID (must be 0-9_999): %d", blockID) + } + objectKey := c.generateS3Key(metadata) - partNumber := blockID + 1 // PartNumber is 1-indexed - _, err := c.s3client.UploadPart(&s3.UploadPartInput{ + partNumber := int32(blockID) + 1 // PartNumber is 1-indexed + _, err := c.s3Client.UploadPart(ctx, &s3.UploadPartInput{ Body: bytes.NewReader(data), Bucket: &c.bucket, Key: &objectKey, @@ -221,15 +234,17 @@ func (c *CarveStore) NewBlock(ctx context.Context, metadata *fleet.CarveMetadata } if blockID >= metadata.BlockCount-1 { // The last block was reached, multipart upload can be completed - parts, err := c.listCompletedParts(objectKey, metadata.SessionId) + parts, err := c.listCompletedParts(ctx, objectKey, metadata.SessionId) if err != nil { return ctxerr.Wrap(ctx, err, "s3 multipart carve upload") } - _, err = c.s3client.CompleteMultipartUpload(&s3.CompleteMultipartUploadInput{ - Bucket: &c.bucket, - Key: &objectKey, - UploadId: &metadata.SessionId, - MultipartUpload: &s3.CompletedMultipartUpload{Parts: parts}, + _, err = c.s3Client.CompleteMultipartUpload(ctx, &s3.CompleteMultipartUploadInput{ + Bucket: &c.bucket, + Key: &objectKey, + UploadId: &metadata.SessionId, + MultipartUpload: &types.CompletedMultipartUpload{ + Parts: parts, + }, }) if err != nil { return ctxerr.Wrap(ctx, err, "s3 multipart carve upload") @@ -246,17 +261,17 @@ func (c *CarveStore) GetBlock(ctx context.Context, metadata *fleet.CarveMetadata // no need to cap the rangeEnd to the carve size as S3 will do that by itself rangeStart := blockID * metadata.BlockSize rangeString := fmt.Sprintf("bytes=%d-%d", rangeStart, rangeStart+metadata.BlockSize-1) - res, err := c.s3client.GetObject(&s3.GetObjectInput{ + res, err := c.s3Client.GetObject(ctx, &s3.GetObjectInput{ Bucket: &c.bucket, Key: &objectKey, Range: &rangeString, }) if err != nil { - var awsErr awserr.Error - if errors.As(err, &awsErr) && awsErr.Code() == s3.ErrCodeNoSuchKey { + var noSuchKey *types.NoSuchKey + if errors.As(err, &noSuchKey) { // The carve does not exists in S3, mark expired metadata.Expired = true - if updateErr := c.UpdateCarve(ctx, metadata); err != nil { + if updateErr := c.UpdateCarve(ctx, metadata); updateErr != nil { err = ctxerr.Wrap(ctx, err, updateErr.Error()) } } diff --git a/server/datastore/s3/common_file_store.go b/server/datastore/s3/common_file_store.go index 15dd023cd2..d48a234f23 100644 --- a/server/datastore/s3/common_file_store.go +++ b/server/datastore/s3/common_file_store.go @@ -10,8 +10,8 @@ import ( "time" "github.com/aws/aws-sdk-go-v2/feature/cloudfront/sign" - "github.com/aws/aws-sdk-go/aws/awserr" - "github.com/aws/aws-sdk-go/service/s3" + "github.com/aws/aws-sdk-go-v2/service/s3" + types "github.com/aws/aws-sdk-go-v2/service/s3/types" "github.com/fleetdm/fleet/v4/server/contexts/ctxerr" "github.com/fleetdm/fleet/v4/server/fleet" "golang.org/x/sync/errgroup" @@ -41,13 +41,18 @@ type commonFileStore struct { func (s *commonFileStore) Get(ctx context.Context, fileID string) (io.ReadCloser, int64, error) { key := s.keyForFile(fileID) - req, err := s.s3client.GetObject(&s3.GetObjectInput{Bucket: &s.bucket, Key: &key}) + req, err := s.s3Client.GetObject(ctx, &s3.GetObjectInput{ + Bucket: &s.bucket, + Key: &key, + }) if err != nil { - if aerr, ok := err.(awserr.Error); ok { - switch aerr.Code() { - case s3.ErrCodeNoSuchKey, s3.ErrCodeNoSuchBucket, "NotFound": - return nil, int64(0), installerNotFoundError{} - } + var ( + noSuchBucket *types.NoSuchBucket + noSuchKey *types.NoSuchKey + notFound *types.NotFound + ) + if errors.As(err, &noSuchBucket) || errors.As(err, &noSuchKey) || errors.As(err, ¬Found) { + return nil, int64(0), installerNotFoundError{} } return nil, int64(0), ctxerr.Wrapf(ctx, err, "retrieving %s from S3 store", s.fileLabel) } @@ -61,7 +66,7 @@ func (s *commonFileStore) Put(ctx context.Context, fileID string, content io.Rea } key := s.keyForFile(fileID) - _, err := s.s3client.PutObject(&s3.PutObjectInput{ + _, err := s.s3Client.PutObject(ctx, &s3.PutObjectInput{ Bucket: &s.bucket, Body: content, Key: &key, @@ -73,13 +78,18 @@ func (s *commonFileStore) Put(ctx context.Context, fileID string, content io.Rea func (s *commonFileStore) Exists(ctx context.Context, fileID string) (bool, error) { key := s.keyForFile(fileID) - _, err := s.s3client.HeadObject(&s3.HeadObjectInput{Bucket: &s.bucket, Key: &key}) + _, err := s.s3Client.HeadObject(ctx, &s3.HeadObjectInput{ + Bucket: &s.bucket, + Key: &key, + }) if err != nil { - if aerr, ok := err.(awserr.Error); ok { - switch aerr.Code() { - case s3.ErrCodeNoSuchKey, s3.ErrCodeNoSuchBucket, "NotFound": - return false, nil - } + var ( + noSuchBucket *types.NoSuchBucket + noSuchKey *types.NoSuchKey + notFound *types.NotFound + ) + if errors.As(err, &noSuchBucket) || errors.As(err, &noSuchKey) || errors.As(err, ¬Found) { + return false, nil } return false, ctxerr.Wrapf(ctx, err, "checking existence of %s in S3 store", s.fileLabel) } @@ -103,7 +113,7 @@ func (s *commonFileStore) Cleanup(ctx context.Context, usedFileIDs []string, rem // again. This approach makes it only two API requests between the read of // used files and the deletions. prefix := path.Join(s.prefix, s.pathPrefix) - page, err := s.s3client.ListObjectsV2(&s3.ListObjectsV2Input{ + page, err := s.s3Client.ListObjectsV2(ctx, &s3.ListObjectsV2Input{ Bucket: &s.bucket, Prefix: &prefix, }) @@ -114,7 +124,7 @@ func (s *commonFileStore) Cleanup(ctx context.Context, usedFileIDs []string, rem // NOTE: there is an inherent risk that we could delete files that were added // between the query to list used IDs and now. We minimize that risk by // checking that the S3 file was created before removeCreatedBefore. - var toDeleteKeys []*s3.ObjectIdentifier + var toDeleteKeys []*types.ObjectIdentifier for _, item := range page.Contents { if item.Key == nil { continue @@ -124,7 +134,7 @@ func (s *commonFileStore) Cleanup(ctx context.Context, usedFileIDs []string, rem } if item.LastModified == nil || !item.LastModified.UTC().After(removeCreatedBefore) { // default to doing the cleanup if we don't have the timestamp information - toDeleteKeys = append(toDeleteKeys, &s3.ObjectIdentifier{Key: item.Key}) + toDeleteKeys = append(toDeleteKeys, &types.ObjectIdentifier{Key: item.Key}) } } @@ -139,7 +149,7 @@ func (s *commonFileStore) Cleanup(ctx context.Context, usedFileIDs []string, rem for _, obj := range toDeleteKeys { obj := obj g.Go(func() error { - _, err := s.s3client.DeleteObject(&s3.DeleteObjectInput{ + _, err := s.s3Client.DeleteObject(ctx, &s3.DeleteObjectInput{ Bucket: &s.bucket, Key: obj.Key, }) diff --git a/server/datastore/s3/s3.go b/server/datastore/s3/s3.go index 321894be34..daee9b4fec 100644 --- a/server/datastore/s3/s3.go +++ b/server/datastore/s3/s3.go @@ -2,23 +2,27 @@ package s3 import ( "context" + "crypto/tls" + "errors" "fmt" - "github.com/aws/aws-sdk-go/aws" - "github.com/aws/aws-sdk-go/aws/awserr" - "github.com/aws/aws-sdk-go/aws/credentials" - "github.com/aws/aws-sdk-go/aws/credentials/stscreds" - "github.com/aws/aws-sdk-go/aws/session" - "github.com/aws/aws-sdk-go/service/s3" - "github.com/aws/aws-sdk-go/service/s3/s3manager" + "github.com/fleetdm/fleet/v4/pkg/fleethttp" "github.com/fleetdm/fleet/v4/server/config" "github.com/fleetdm/fleet/v4/server/fleet" + + "github.com/aws/aws-sdk-go-v2/aws" + aws_config "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/credentials" + "github.com/aws/aws-sdk-go-v2/credentials/stscreds" + "github.com/aws/aws-sdk-go-v2/feature/s3/manager" + "github.com/aws/aws-sdk-go-v2/service/s3" + types "github.com/aws/aws-sdk-go-v2/service/s3/types" ) const awsRegionHint = "us-east-1" type s3store struct { - s3client *s3.S3 + s3Client *s3.Client bucket string prefix string cloudFrontConfig *config.S3CloudFrontConfig @@ -36,76 +40,99 @@ func (p installerNotFoundError) IsNotFound() bool { return true } -// newS3store initializes an S3 Datastore -func newS3store(config config.S3ConfigInternal) (*s3store, error) { - conf := &aws.Config{} +// newS3Store initializes an S3 Datastore. +func newS3Store(cfg config.S3ConfigInternal) (*s3store, error) { + var opts []func(*aws_config.LoadOptions) error - // Use default auth provire if no static credentials were provided - if config.AccessKeyID != "" && config.SecretAccessKey != "" { - conf.Credentials = credentials.NewStaticCredentials( - config.AccessKeyID, - config.SecretAccessKey, - "", + // The service endpoint is deprecated, but we still set it + // in case users are using it. + // It is also used when testing with minio. + if cfg.EndpointURL != "" { + opts = append(opts, aws_config.WithEndpointResolver(aws.EndpointResolverFunc( + func(service, region string) (aws.Endpoint, error) { + return aws.Endpoint{ + URL: cfg.EndpointURL, + }, nil + })), ) } - if config.EndpointURL != "" { - conf.Endpoint = &config.EndpointURL + // DisableSSL is only used for testing. + if cfg.DisableSSL { + // Ignoring "G402: TLS InsecureSkipVerify set true", this is only used for automated testing. + c := fleethttp.NewClient(fleethttp.WithTLSClientConfig(&tls.Config{ //nolint:gosec + InsecureSkipVerify: false, + })) + opts = append(opts, aws_config.WithHTTPClient(c)) } - conf.DisableSSL = &config.DisableSSL - conf.S3ForcePathStyle = &config.ForceS3PathStyle - - sess, err := session.NewSession(conf) - if err != nil { - return nil, fmt.Errorf("create S3 client: %w", err) + // Use default auth provider if no static credentials were provided. + if cfg.AccessKeyID != "" && cfg.SecretAccessKey != "" { + opts = append(opts, aws_config.WithCredentialsProvider(credentials.NewStaticCredentialsProvider( + cfg.AccessKeyID, + cfg.SecretAccessKey, + "", + ))) } - // Assume role if configured - if config.StsAssumeRoleArn != "" { - creds := stscreds.NewCredentials(sess, config.StsAssumeRoleArn, func(provider *stscreds.AssumeRoleProvider) { - if config.StsExternalID != "" { - provider.ExternalID = &config.StsExternalID + // cfg.StsAssumeRoleArn has been marked as deprecated, but we still set it in case users are using it. + if cfg.StsAssumeRoleArn != "" { + opts = append(opts, aws_config.WithAssumeRoleCredentialOptions(func(r *stscreds.AssumeRoleOptions) { + r.RoleARN = cfg.StsAssumeRoleArn + if cfg.StsExternalID != "" { + r.ExternalID = &cfg.StsExternalID } - }) - conf.Credentials = creds - sess, err = session.NewSession(conf) - if err != nil { - return nil, fmt.Errorf("create S3 client: %w", err) - } + })) } - if len(config.Region) == 0 { - region, err := s3manager.GetBucketRegion(context.TODO(), sess, config.Bucket, awsRegionHint) + if cfg.Region == "" { + // Attempt to deduce region from bucket. + conf, err := aws_config.LoadDefaultConfig(context.Background(), + append(opts, aws_config.WithRegion(awsRegionHint))..., + ) if err != nil { - return nil, fmt.Errorf("create S3 client: %w", err) + return nil, fmt.Errorf("failed to create default config to get bucket region: %w", err) } - config.Region = region + bucketRegion, err := manager.GetBucketRegion(context.Background(), s3.NewFromConfig(conf), cfg.Bucket) + if err != nil { + return nil, fmt.Errorf("get bucket region: %w", err) + } + cfg.Region = bucketRegion } + opts = append(opts, aws_config.WithRegion(cfg.Region)) + conf, err := aws_config.LoadDefaultConfig(context.Background(), opts...) + if err != nil { + return nil, fmt.Errorf("failed to create default config: %w", err) + } + + s3Client := s3.NewFromConfig(conf, func(o *s3.Options) { + o.UsePathStyle = cfg.ForceS3PathStyle + }) + return &s3store{ - s3client: s3.New(sess, &aws.Config{Region: &config.Region}), - bucket: config.Bucket, - prefix: config.Prefix, - cloudFrontConfig: config.CloudFrontConfig, + s3Client: s3Client, + bucket: cfg.Bucket, + prefix: cfg.Prefix, + cloudFrontConfig: cfg.CloudFrontConfig, }, nil } // CreateTestBucket creates a bucket with the provided name and a default // bucket config. Only recommended for local testing. -func (s *s3store) CreateTestBucket(name string) error { - _, err := s.s3client.CreateBucket(&s3.CreateBucketInput{ +func (s *s3store) CreateTestBucket(ctx context.Context, name string) error { + _, err := s.s3Client.CreateBucket(ctx, &s3.CreateBucketInput{ Bucket: &name, - CreateBucketConfiguration: &s3.CreateBucketConfiguration{}, + CreateBucketConfiguration: &types.CreateBucketConfiguration{}, }) // Don't error if the bucket already exists - if aerr, ok := err.(awserr.Error); ok { - switch aerr.Code() { - case s3.ErrCodeBucketAlreadyExists, s3.ErrCodeBucketAlreadyOwnedByYou: - return nil - } + var ( + bucketAlreadyExists *types.BucketAlreadyExists + bucketAlreadyOwnedByYou *types.BucketAlreadyOwnedByYou + ) + if errors.As(err, &bucketAlreadyExists) || errors.As(err, &bucketAlreadyOwnedByYou) { + return nil } - return err } diff --git a/server/datastore/s3/software_installer.go b/server/datastore/s3/software_installer.go index 60494d83f4..e64751cccd 100644 --- a/server/datastore/s3/software_installer.go +++ b/server/datastore/s3/software_installer.go @@ -12,7 +12,7 @@ type SoftwareInstallerStore struct { // NewSoftwareInstallerStore creates a new instance with the given S3 config. func NewSoftwareInstallerStore(config config.S3Config) (*SoftwareInstallerStore, error) { - s3store, err := newS3store(config.SoftwareInstallersToInternalCfg()) + s3store, err := newS3Store(config.SoftwareInstallersToInternalCfg()) if err != nil { return nil, err } diff --git a/server/datastore/s3/software_installer_test.go b/server/datastore/s3/software_installer_test.go index f049121468..9be28e3916 100644 --- a/server/datastore/s3/software_installer_test.go +++ b/server/datastore/s3/software_installer_test.go @@ -12,7 +12,7 @@ import ( "testing" "time" - "github.com/aws/aws-sdk-go/service/s3" + "github.com/aws/aws-sdk-go-v2/service/s3" "github.com/fleetdm/fleet/v4/server/fleet" "github.com/google/uuid" "github.com/stretchr/testify/require" @@ -91,7 +91,7 @@ func TestSoftwareInstallerCleanup(t *testing.T) { assertExisting := func(want []string) { prefix := path.Join(store.prefix, softwareInstallersPrefix) - page, err := store.s3client.ListObjectsV2(&s3.ListObjectsV2Input{ + page, err := store.s3Client.ListObjectsV2(ctx, &s3.ListObjectsV2Input{ Bucket: &store.bucket, Prefix: &prefix, }) diff --git a/server/datastore/s3/testing_utils.go b/server/datastore/s3/testing_utils.go index 921b9da688..ca0062e27f 100644 --- a/server/datastore/s3/testing_utils.go +++ b/server/datastore/s3/testing_utils.go @@ -1,11 +1,13 @@ package s3 import ( + "context" + "errors" "os" "testing" - "github.com/aws/aws-sdk-go/aws/awserr" - "github.com/aws/aws-sdk-go/service/s3" + "github.com/aws/aws-sdk-go-v2/service/s3" + "github.com/aws/aws-sdk-go-v2/service/s3/types" "github.com/fleetdm/fleet/v4/server/config" "github.com/stretchr/testify/require" ) @@ -13,7 +15,7 @@ import ( const ( accessKeyID = "minio" secretAccessKey = "minio123!" - testEndpoint = "localhost:9000" + testEndpoint = "http://localhost:9000" ) func SetupTestSoftwareInstallerStore(tb testing.TB, bucket, prefix string) *SoftwareInstallerStore { @@ -29,7 +31,7 @@ func SetupTestBootstrapPackageStore(tb testing.TB, bucket, prefix string) *Boots } type testBucketCreator interface { - CreateTestBucket(name string) error + CreateTestBucket(ctx context.Context, name string) error } func setupTestStore[T testBucketCreator](tb testing.TB, bucket, prefix string, newFn func(config.S3Config) (T, error)) T { @@ -56,7 +58,7 @@ func setupTestStore[T testBucketCreator](tb testing.TB, bucket, prefix string, n }) require.Nil(tb, err) - err = store.CreateTestBucket(bucket) + err = store.CreateTestBucket(context.Background(), bucket) require.NoError(tb, err) return store @@ -65,32 +67,34 @@ func setupTestStore[T testBucketCreator](tb testing.TB, bucket, prefix string, n func cleanupStore(tb testing.TB, store *s3store) { checkEnv(tb) - resp, err := store.s3client.ListObjects(&s3.ListObjectsInput{ + ctx := context.Background() + resp, err := store.s3Client.ListObjects(ctx, &s3.ListObjectsInput{ Bucket: &store.bucket, }) - if aerr, ok := err.(awserr.Error); ok { - if aerr.Code() == s3.ErrCodeNoSuchBucket { - // fine, nothing to clean-up if the bucket no longer exists, no error - return - } + var noSuchBucket *types.NoSuchBucket + if errors.As(err, &noSuchBucket) { + // OK, nothing to clean-up if the bucket no longer exists, no error + return } require.NoError(tb, err) - var objs []*s3.ObjectIdentifier + var objs []types.ObjectIdentifier for _, o := range resp.Contents { - objs = append(objs, &s3.ObjectIdentifier{Key: o.Key}) + objs = append(objs, types.ObjectIdentifier{ + Key: o.Key, + }) } if len(objs) > 0 { - _, err = store.s3client.DeleteObjects(&s3.DeleteObjectsInput{ + _, err = store.s3Client.DeleteObjects(ctx, &s3.DeleteObjectsInput{ Bucket: &store.bucket, - Delete: &s3.Delete{ + Delete: &types.Delete{ Objects: objs, }, }) require.NoError(tb, err) } - _, err = store.s3client.DeleteBucket(&s3.DeleteBucketInput{ + _, err = store.s3Client.DeleteBucket(ctx, &s3.DeleteBucketInput{ Bucket: &store.bucket, }) require.NoError(tb, err) diff --git a/server/fleet/emails.go b/server/fleet/emails.go index 941dc5b9ef..1ec2d5c485 100644 --- a/server/fleet/emails.go +++ b/server/fleet/emails.go @@ -1,6 +1,7 @@ package fleet import ( + "context" "errors" "regexp" "time" @@ -24,7 +25,7 @@ type Email struct { } type MailService interface { - SendEmail(e Email) error + SendEmail(ctx context.Context, e Email) error CanSendEmail(smtpSettings SMTPSettings) bool } diff --git a/server/logging/firehose.go b/server/logging/firehose.go index 4e0d371852..1528c73d88 100644 --- a/server/logging/firehose.go +++ b/server/logging/firehose.go @@ -8,13 +8,12 @@ import ( "math" "time" - "github.com/aws/aws-sdk-go/aws" - "github.com/aws/aws-sdk-go/aws/awserr" - "github.com/aws/aws-sdk-go/aws/credentials" - "github.com/aws/aws-sdk-go/aws/credentials/stscreds" - "github.com/aws/aws-sdk-go/aws/session" - "github.com/aws/aws-sdk-go/service/firehose" - "github.com/aws/aws-sdk-go/service/firehose/firehoseiface" + "github.com/aws/aws-sdk-go-v2/aws" + aws_config "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/credentials" + "github.com/aws/aws-sdk-go-v2/credentials/stscreds" + "github.com/aws/aws-sdk-go-v2/service/firehose" + "github.com/aws/aws-sdk-go-v2/service/firehose/types" "github.com/fleetdm/fleet/v4/server/contexts/ctxerr" "github.com/go-kit/log" "github.com/go-kit/log/level" @@ -31,58 +30,70 @@ const ( firehoseMaxSizeOfBatch = 4 * 1000 * 1000 // 4 MB ) +type FirehoseAPI interface { + PutRecordBatch(ctx context.Context, input *firehose.PutRecordBatchInput, optFns ...func(*firehose.Options)) (*firehose.PutRecordBatchOutput, error) + DescribeDeliveryStream(ctx context.Context, input *firehose.DescribeDeliveryStreamInput, optFns ...func(*firehose.Options)) (*firehose.DescribeDeliveryStreamOutput, error) +} + type firehoseLogWriter struct { - client firehoseiface.FirehoseAPI + client FirehoseAPI stream string logger log.Logger } func NewFirehoseLogWriter(region, endpointURL, id, secret, stsAssumeRoleArn, stsExternalID, stream string, logger log.Logger) (*firehoseLogWriter, error) { - conf := &aws.Config{ - Region: ®ion, - Endpoint: &endpointURL, // empty string or nil will use default values + var opts []func(*aws_config.LoadOptions) error + + // The service endpoint is deprecated, but we still set it + // in case users are using it. + if endpointURL != "" { + opts = append(opts, aws_config.WithEndpointResolver(aws.EndpointResolverFunc( + func(service, region string) (aws.Endpoint, error) { + return aws.Endpoint{ + URL: endpointURL, + }, nil + })), + ) } // Only provide static credentials if we have them - // otherwise use the default credentials provider chain + // otherwise use the default credentials provider chain. if id != "" && secret != "" { - conf.Credentials = credentials.NewStaticCredentials(id, secret, "") - } - - sess, err := session.NewSession(conf) - if err != nil { - return nil, fmt.Errorf("create Firehose client: %w", err) + opts = append(opts, + aws_config.WithCredentialsProvider(credentials.NewStaticCredentialsProvider(id, secret, "")), + ) } if stsAssumeRoleArn != "" { - creds := stscreds.NewCredentials(sess, stsAssumeRoleArn, func(provider *stscreds.AssumeRoleProvider) { + opts = append(opts, aws_config.WithAssumeRoleCredentialOptions(func(r *stscreds.AssumeRoleOptions) { + r.RoleARN = stsAssumeRoleArn if stsExternalID != "" { - provider.ExternalID = &stsExternalID + r.ExternalID = &stsExternalID } - }) - conf.Credentials = creds - - sess, err = session.NewSession(conf) - - if err != nil { - return nil, fmt.Errorf("create Firehose client: %w", err) - } + })) } - client := firehose.New(sess) + + opts = append(opts, aws_config.WithRegion(region)) + conf, err := aws_config.LoadDefaultConfig(context.Background(), opts...) + if err != nil { + return nil, fmt.Errorf("failed to create default config: %w", err) + } + + firehoseClient := firehose.NewFromConfig(conf) f := &firehoseLogWriter{ - client: client, + client: firehoseClient, stream: stream, logger: logger, } - if err := f.validateStream(); err != nil { - return nil, fmt.Errorf("create Firehose writer: %w", err) + if err := f.validateStream(context.Background()); err != nil { + return nil, fmt.Errorf("validate firehose: %w", err) } return f, nil } -func (f *firehoseLogWriter) validateStream() error { - out, err := f.client.DescribeDeliveryStream( +func (f *firehoseLogWriter) validateStream(ctx context.Context) error { + out, err := f.client.DescribeDeliveryStream(ctx, &firehose.DescribeDeliveryStreamInput{ DeliveryStreamName: &f.stream, }, @@ -91,7 +102,7 @@ func (f *firehoseLogWriter) validateStream() error { return fmt.Errorf("describe stream %s: %w", f.stream, err) } - if (*out.DeliveryStreamDescription.DeliveryStreamStatus) != firehose.DeliveryStreamStatusActive { + if out.DeliveryStreamDescription.DeliveryStreamStatus != types.DeliveryStreamStatusActive { return fmt.Errorf("delivery stream %s not active", f.stream) } @@ -99,7 +110,7 @@ func (f *firehoseLogWriter) validateStream() error { } func (f *firehoseLogWriter) Write(ctx context.Context, logs []json.RawMessage) error { - var records []*firehose.Record + var records []types.Record totalBytes := 0 for _, log := range logs { // Add newline because Firehose does not output each record on @@ -126,20 +137,22 @@ func (f *firehoseLogWriter) Write(ctx context.Context, logs []json.RawMessage) e // adding any more. if len(records) >= firehoseMaxRecordsInBatch || totalBytes+len(log) > firehoseMaxSizeOfBatch { - if err := f.putRecordBatch(0, records); err != nil { + if err := f.putRecordBatch(ctx, 0, records); err != nil { return ctxerr.Wrap(ctx, err, "put records") } totalBytes = 0 records = nil } - records = append(records, &firehose.Record{Data: []byte(log)}) + records = append(records, types.Record{ + Data: []byte(log), + }) totalBytes += len(log) } // Push the final batch if len(records) > 0 { - if err := f.putRecordBatch(0, records); err != nil { + if err := f.putRecordBatch(ctx, 0, records); err != nil { return ctxerr.Wrap(ctx, err, "put records") } } @@ -147,7 +160,7 @@ func (f *firehoseLogWriter) Write(ctx context.Context, logs []json.RawMessage) e return nil } -func (f *firehoseLogWriter) putRecordBatch(try int, records []*firehose.Record) error { +func (f *firehoseLogWriter) putRecordBatch(ctx context.Context, try int, records []types.Record) error { if try > 0 { time.Sleep(100 * time.Millisecond * time.Duration(math.Pow(2.0, float64(try)))) } @@ -156,13 +169,13 @@ func (f *firehoseLogWriter) putRecordBatch(try int, records []*firehose.Record) Records: records, } - output, err := f.client.PutRecordBatch(input) + output, err := f.client.PutRecordBatch(ctx, input) if err != nil { - var aerr awserr.Error - if errors.As(err, &aerr) { - if aerr.Code() == firehose.ErrCodeServiceUnavailableException && try < firehoseMaxRetries { + var serviceUnavailableException *types.ServiceUnavailableException + if errors.As(err, &serviceUnavailableException) { + if try < firehoseMaxRetries { // Retry with backoff - return f.putRecordBatch(try+1, records) + return f.putRecordBatch(ctx, try+1, records) } } @@ -190,7 +203,7 @@ func (f *firehoseLogWriter) putRecordBatch(try int, records []*firehose.Record) ) } - var failedRecords []*firehose.Record + var failedRecords []types.Record // Collect failed records for retry for i, record := range output.RequestResponses { if record.ErrorCode != nil { @@ -198,7 +211,7 @@ func (f *firehoseLogWriter) putRecordBatch(try int, records []*firehose.Record) } } - return f.putRecordBatch(try+1, failedRecords) + return f.putRecordBatch(ctx, try+1, failedRecords) } return nil diff --git a/server/logging/firehose_test.go b/server/logging/firehose_test.go index 667727f61d..390092f96b 100644 --- a/server/logging/firehose_test.go +++ b/server/logging/firehose_test.go @@ -6,10 +6,9 @@ import ( "errors" "testing" - "github.com/aws/aws-sdk-go/aws" - "github.com/aws/aws-sdk-go/aws/awserr" - "github.com/aws/aws-sdk-go/service/firehose" - "github.com/aws/aws-sdk-go/service/firehose/firehoseiface" + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/firehose" + "github.com/aws/aws-sdk-go-v2/service/firehose/types" "github.com/fleetdm/fleet/v4/server/logging/mock" "github.com/go-kit/log" "github.com/stretchr/testify/assert" @@ -28,7 +27,7 @@ var ( } ) -func makeFirehoseWriterWithMock(client firehoseiface.FirehoseAPI, stream string) *firehoseLogWriter { +func makeFirehoseWriterWithMock(client FirehoseAPI, stream string) *firehoseLogWriter { return &firehoseLogWriter{ client: client, stream: stream, @@ -47,7 +46,7 @@ func getLogsFromInput(input *firehose.PutRecordBatchInput) []json.RawMessage { func TestFirehoseNonRetryableFailure(t *testing.T) { ctx := context.Background() callCount := 0 - putFunc := func(*firehose.PutRecordBatchInput) (*firehose.PutRecordBatchOutput, error) { + putFunc := func(context.Context, *firehose.PutRecordBatchInput, ...func(*firehose.Options)) (*firehose.PutRecordBatchOutput, error) { callCount += 1 return nil, errors.New("generic error") } @@ -61,12 +60,12 @@ func TestFirehoseNonRetryableFailure(t *testing.T) { func TestFirehoseRetryableFailure(t *testing.T) { ctx := context.Background() callCount := 0 - putFunc := func(input *firehose.PutRecordBatchInput) (*firehose.PutRecordBatchOutput, error) { + putFunc := func(ctx context.Context, input *firehose.PutRecordBatchInput, optFns ...func(*firehose.Options)) (*firehose.PutRecordBatchOutput, error) { callCount += 1 assert.Equal(t, logsWithNewlines, getLogsFromInput(input)) assert.Equal(t, "foobar", *input.DeliveryStreamName) if callCount < 3 { - return nil, awserr.New(firehose.ErrCodeServiceUnavailableException, "", nil) + return nil, &types.ServiceUnavailableException{} } // Returning a non-retryable error earlier helps keep // this test faster @@ -82,11 +81,13 @@ func TestFirehoseRetryableFailure(t *testing.T) { func TestFirehoseNormalPut(t *testing.T) { ctx := context.Background() callCount := 0 - putFunc := func(input *firehose.PutRecordBatchInput) (*firehose.PutRecordBatchOutput, error) { + putFunc := func(ctx context.Context, input *firehose.PutRecordBatchInput, optFns ...func(*firehose.Options)) (*firehose.PutRecordBatchOutput, error) { callCount += 1 assert.Equal(t, logsWithNewlines, getLogsFromInput(input)) assert.Equal(t, "foobar", *input.DeliveryStreamName) - return &firehose.PutRecordBatchOutput{FailedPutCount: aws.Int64(0)}, nil + return &firehose.PutRecordBatchOutput{ + FailedPutCount: aws.Int32(0), + }, nil } f := &mock.FirehoseMock{PutRecordBatchFunc: putFunc} writer := makeFirehoseWriterWithMock(f, "foobar") @@ -100,48 +101,48 @@ func TestFirehoseSomeFailures(t *testing.T) { f := &mock.FirehoseMock{} callCount := 0 - call3 := func(input *firehose.PutRecordBatchInput) (*firehose.PutRecordBatchOutput, error) { + call3 := func(ctx context.Context, input *firehose.PutRecordBatchInput, optFns ...func(*firehose.Options)) (*firehose.PutRecordBatchOutput, error) { // final invocation callCount += 1 assert.Equal(t, logsWithNewlines[1:2], getLogsFromInput(input)) return &firehose.PutRecordBatchOutput{ - FailedPutCount: aws.Int64(0), + FailedPutCount: aws.Int32(0), }, nil } - call2 := func(input *firehose.PutRecordBatchInput) (*firehose.PutRecordBatchOutput, error) { + call2 := func(ctx context.Context, input *firehose.PutRecordBatchInput, optFns ...func(*firehose.Options)) (*firehose.PutRecordBatchOutput, error) { // Set to invoke call3 next time f.PutRecordBatchFunc = call3 callCount += 1 assert.Equal(t, logsWithNewlines[1:], getLogsFromInput(input)) return &firehose.PutRecordBatchOutput{ - FailedPutCount: aws.Int64(1), - RequestResponses: []*firehose.PutRecordBatchResponseEntry{ - &firehose.PutRecordBatchResponseEntry{ + FailedPutCount: aws.Int32(1), + RequestResponses: []types.PutRecordBatchResponseEntry{ + { ErrorCode: aws.String("error"), }, - &firehose.PutRecordBatchResponseEntry{ + { RecordId: aws.String("foo"), }, }, }, nil } - call1 := func(input *firehose.PutRecordBatchInput) (*firehose.PutRecordBatchOutput, error) { + call1 := func(ctx context.Context, input *firehose.PutRecordBatchInput, optFns ...func(*firehose.Options)) (*firehose.PutRecordBatchOutput, error) { // Use call2 function for next call f.PutRecordBatchFunc = call2 callCount += 1 assert.Equal(t, logsWithNewlines, getLogsFromInput(input)) return &firehose.PutRecordBatchOutput{ - FailedPutCount: aws.Int64(1), - RequestResponses: []*firehose.PutRecordBatchResponseEntry{ - &firehose.PutRecordBatchResponseEntry{ + FailedPutCount: aws.Int32(1), + RequestResponses: []types.PutRecordBatchResponseEntry{ + { RecordId: aws.String("foo"), }, - &firehose.PutRecordBatchResponseEntry{ + { ErrorCode: aws.String("error"), }, - &firehose.PutRecordBatchResponseEntry{ + { ErrorCode: aws.String("error"), }, }, @@ -159,13 +160,13 @@ func TestFirehoseFailAllRecords(t *testing.T) { f := &mock.FirehoseMock{} callCount := 0 - f.PutRecordBatchFunc = func(input *firehose.PutRecordBatchInput) (*firehose.PutRecordBatchOutput, error) { + f.PutRecordBatchFunc = func(ctx context.Context, input *firehose.PutRecordBatchInput, optFns ...func(*firehose.Options)) (*firehose.PutRecordBatchOutput, error) { callCount += 1 assert.Equal(t, logsWithNewlines, getLogsFromInput(input)) if callCount < 3 { return &firehose.PutRecordBatchOutput{ - FailedPutCount: aws.Int64(1), - RequestResponses: []*firehose.PutRecordBatchResponseEntry{ + FailedPutCount: aws.Int32(1), + RequestResponses: []types.PutRecordBatchResponseEntry{ {ErrorCode: aws.String("error")}, {ErrorCode: aws.String("error")}, {ErrorCode: aws.String("error")}, @@ -189,11 +190,11 @@ func TestFirehoseRecordTooBig(t *testing.T) { copy(newLogs, logs) newLogs[0] = make(json.RawMessage, firehoseMaxSizeOfRecord+1) callCount := 0 - putFunc := func(input *firehose.PutRecordBatchInput) (*firehose.PutRecordBatchOutput, error) { + putFunc := func(ctx context.Context, input *firehose.PutRecordBatchInput, optFns ...func(*firehose.Options)) (*firehose.PutRecordBatchOutput, error) { callCount += 1 assert.Equal(t, logsWithNewlines[1:], getLogsFromInput(input)) assert.Equal(t, "foobar", *input.DeliveryStreamName) - return &firehose.PutRecordBatchOutput{FailedPutCount: aws.Int64(0)}, nil + return &firehose.PutRecordBatchOutput{FailedPutCount: aws.Int32(0)}, nil } f := &mock.FirehoseMock{PutRecordBatchFunc: putFunc} writer := makeFirehoseWriterWithMock(f, "foobar") @@ -211,11 +212,11 @@ func TestFirehoseSplitBatchBySize(t *testing.T) { logs[i] = make(json.RawMessage, firehoseMaxSizeOfRecord-1) } callCount := 0 - putFunc := func(input *firehose.PutRecordBatchInput) (*firehose.PutRecordBatchOutput, error) { + putFunc := func(ctx context.Context, input *firehose.PutRecordBatchInput, optFns ...func(*firehose.Options)) (*firehose.PutRecordBatchOutput, error) { callCount += 1 assert.Len(t, getLogsFromInput(input), 4) assert.Equal(t, "foobar", *input.DeliveryStreamName) - return &firehose.PutRecordBatchOutput{FailedPutCount: aws.Int64(0)}, nil + return &firehose.PutRecordBatchOutput{FailedPutCount: aws.Int32(0)}, nil } f := &mock.FirehoseMock{PutRecordBatchFunc: putFunc} writer := makeFirehoseWriterWithMock(f, "foobar") @@ -231,11 +232,11 @@ func TestFirehoseSplitBatchByCount(t *testing.T) { logs[i] = json.RawMessage(`{}`) } callCount := 0 - putFunc := func(input *firehose.PutRecordBatchInput) (*firehose.PutRecordBatchOutput, error) { + putFunc := func(ctx context.Context, input *firehose.PutRecordBatchInput, optFns ...func(*firehose.Options)) (*firehose.PutRecordBatchOutput, error) { callCount += 1 assert.Len(t, getLogsFromInput(input), 500) assert.Equal(t, "foobar", *input.DeliveryStreamName) - return &firehose.PutRecordBatchOutput{FailedPutCount: aws.Int64(0)}, nil + return &firehose.PutRecordBatchOutput{FailedPutCount: aws.Int32(0)}, nil } f := &mock.FirehoseMock{PutRecordBatchFunc: putFunc} writer := makeFirehoseWriterWithMock(f, "foobar") @@ -245,45 +246,45 @@ func TestFirehoseSplitBatchByCount(t *testing.T) { } func TestFirehoseValidateStreamActive(t *testing.T) { - describeFunc := func(input *firehose.DescribeDeliveryStreamInput) (*firehose.DescribeDeliveryStreamOutput, error) { + describeFunc := func(ctx context.Context, input *firehose.DescribeDeliveryStreamInput, optFns ...func(*firehose.Options)) (*firehose.DescribeDeliveryStreamOutput, error) { assert.Equal(t, "test", *input.DeliveryStreamName) return &firehose.DescribeDeliveryStreamOutput{ - DeliveryStreamDescription: &firehose.DeliveryStreamDescription{ - DeliveryStreamStatus: aws.String(firehose.DeliveryStreamStatusActive), + DeliveryStreamDescription: &types.DeliveryStreamDescription{ + DeliveryStreamStatus: types.DeliveryStreamStatusActive, }, }, nil } f := &mock.FirehoseMock{DescribeDeliveryStreamFunc: describeFunc} writer := makeFirehoseWriterWithMock(f, "test") - err := writer.validateStream() + err := writer.validateStream(context.Background()) assert.NoError(t, err) assert.True(t, f.DescribeDeliveryStreamFuncInvoked) } func TestFirehoseValidateStreamNotActive(t *testing.T) { - describeFunc := func(input *firehose.DescribeDeliveryStreamInput) (*firehose.DescribeDeliveryStreamOutput, error) { + describeFunc := func(ctx context.Context, input *firehose.DescribeDeliveryStreamInput, optFns ...func(*firehose.Options)) (*firehose.DescribeDeliveryStreamOutput, error) { assert.Equal(t, "test", *input.DeliveryStreamName) return &firehose.DescribeDeliveryStreamOutput{ - DeliveryStreamDescription: &firehose.DeliveryStreamDescription{ - DeliveryStreamStatus: aws.String(firehose.DeliveryStreamStatusCreating), + DeliveryStreamDescription: &types.DeliveryStreamDescription{ + DeliveryStreamStatus: types.DeliveryStreamStatusCreating, }, }, nil } f := &mock.FirehoseMock{DescribeDeliveryStreamFunc: describeFunc} writer := makeFirehoseWriterWithMock(f, "test") - err := writer.validateStream() + err := writer.validateStream(context.Background()) assert.Error(t, err) assert.True(t, f.DescribeDeliveryStreamFuncInvoked) } func TestFirehoseValidateStreamError(t *testing.T) { - describeFunc := func(input *firehose.DescribeDeliveryStreamInput) (*firehose.DescribeDeliveryStreamOutput, error) { + describeFunc := func(ctx context.Context, input *firehose.DescribeDeliveryStreamInput, optFns ...func(*firehose.Options)) (*firehose.DescribeDeliveryStreamOutput, error) { assert.Equal(t, "test", *input.DeliveryStreamName) return nil, errors.New("boom!") } f := &mock.FirehoseMock{DescribeDeliveryStreamFunc: describeFunc} writer := makeFirehoseWriterWithMock(f, "test") - err := writer.validateStream() + err := writer.validateStream(context.Background()) assert.Error(t, err) assert.True(t, f.DescribeDeliveryStreamFuncInvoked) } diff --git a/server/logging/kinesis.go b/server/logging/kinesis.go index 717662c639..f018aa6848 100644 --- a/server/logging/kinesis.go +++ b/server/logging/kinesis.go @@ -9,13 +9,13 @@ import ( "math/rand" "time" - "github.com/aws/aws-sdk-go/aws" - "github.com/aws/aws-sdk-go/aws/awserr" - "github.com/aws/aws-sdk-go/aws/credentials" - "github.com/aws/aws-sdk-go/aws/credentials/stscreds" - "github.com/aws/aws-sdk-go/aws/session" - "github.com/aws/aws-sdk-go/service/kinesis" - "github.com/aws/aws-sdk-go/service/kinesis/kinesisiface" + "github.com/aws/aws-sdk-go-v2/aws" + aws_config "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/credentials" + "github.com/aws/aws-sdk-go-v2/credentials/stscreds" + "github.com/aws/aws-sdk-go-v2/service/kinesis" + "github.com/aws/aws-sdk-go-v2/service/kinesis/types" + smithy "github.com/aws/smithy-go" "github.com/fleetdm/fleet/v4/server/contexts/ctxerr" "github.com/go-kit/log" "github.com/go-kit/log/level" @@ -32,64 +32,77 @@ const ( kinesisMaxSizeOfBatch = 5 * 1000 * 1000 // 5 MB ) +type KinesisAPI interface { + PutRecords(ctx context.Context, params *kinesis.PutRecordsInput, optFns ...func(*kinesis.Options)) (*kinesis.PutRecordsOutput, error) + DescribeStream(ctx context.Context, params *kinesis.DescribeStreamInput, optFns ...func(*kinesis.Options)) (*kinesis.DescribeStreamOutput, error) +} + type kinesisLogWriter struct { - client kinesisiface.KinesisAPI + client KinesisAPI stream string logger log.Logger rand *rand.Rand } func NewKinesisLogWriter(region, endpointURL, id, secret, stsAssumeRoleArn, stsExternalID, stream string, logger log.Logger) (*kinesisLogWriter, error) { - conf := &aws.Config{ - Region: ®ion, - Endpoint: &endpointURL, // empty string or nil will use default values + var opts []func(*aws_config.LoadOptions) error + + // The service endpoint is deprecated, but we still set it + // in case users are using it. + if endpointURL != "" { + opts = append(opts, aws_config.WithEndpointResolver(aws.EndpointResolverFunc( + func(service, region string) (aws.Endpoint, error) { + return aws.Endpoint{ + URL: endpointURL, + }, nil + })), + ) } // Only provide static credentials if we have them - // otherwise use the default credentials provider chain + // otherwise use the default credentials provider chain. if id != "" && secret != "" { - conf.Credentials = credentials.NewStaticCredentials(id, secret, "") - } - - sess, err := session.NewSession(conf) - if err != nil { - return nil, fmt.Errorf("create Kinesis client: %w", err) + opts = append(opts, + aws_config.WithCredentialsProvider(credentials.NewStaticCredentialsProvider(id, secret, "")), + ) } + // cfg.StsAssumeRoleArn has been marked as deprecated, but we still set it in case users are using it. if stsAssumeRoleArn != "" { - creds := stscreds.NewCredentials(sess, stsAssumeRoleArn, func(provider *stscreds.AssumeRoleProvider) { + opts = append(opts, aws_config.WithAssumeRoleCredentialOptions(func(r *stscreds.AssumeRoleOptions) { + r.RoleARN = stsAssumeRoleArn if stsExternalID != "" { - provider.ExternalID = &stsExternalID + r.ExternalID = &stsExternalID } - }) - conf.Credentials = creds - - sess, err = session.NewSession(conf) - - if err != nil { - return nil, fmt.Errorf("create Kinesis client: %w", err) - } + })) } - client := kinesis.New(sess) + + opts = append(opts, aws_config.WithRegion(region)) + conf, err := aws_config.LoadDefaultConfig(context.Background(), opts...) + if err != nil { + return nil, fmt.Errorf("failed to create default config: %w", err) + } + + kinesisClient := kinesis.NewFromConfig(conf) // This will be used to generate random partition keys to balance // records across Kinesis shards. rand := rand.New(rand.NewSource(time.Now().UnixNano())) k := &kinesisLogWriter{ - client: client, + client: kinesisClient, stream: stream, logger: logger, rand: rand, } - if err := k.validateStream(); err != nil { - return nil, fmt.Errorf("create Kinesis writer: %w", err) + if err := k.validateStream(context.Background()); err != nil { + return nil, fmt.Errorf("validate kinesis: %w", err) } return k, nil } -func (k *kinesisLogWriter) validateStream() error { - out, err := k.client.DescribeStream( +func (k *kinesisLogWriter) validateStream(ctx context.Context) error { + out, err := k.client.DescribeStream(ctx, &kinesis.DescribeStreamInput{ StreamName: &k.stream, }, @@ -98,7 +111,7 @@ func (k *kinesisLogWriter) validateStream() error { return fmt.Errorf("describe stream %s: %w", k.stream, err) } - if (*out.StreamDescription.StreamStatus) != kinesis.StreamStatusActive { + if out.StreamDescription.StreamStatus != types.StreamStatusActive { return fmt.Errorf("stream %s not active", k.stream) } @@ -106,7 +119,7 @@ func (k *kinesisLogWriter) validateStream() error { } func (k *kinesisLogWriter) Write(ctx context.Context, logs []json.RawMessage) error { - var records []*kinesis.PutRecordsRequestEntry + var records []types.PutRecordsRequestEntry totalBytes := 0 for _, log := range logs { // so we get nice NDJSON @@ -135,20 +148,23 @@ func (k *kinesisLogWriter) Write(ctx context.Context, logs []json.RawMessage) er // adding any more. if len(records) >= kinesisMaxRecordsInBatch || totalBytes+len(log)+len(partitionKey) > kinesisMaxSizeOfBatch { - if err := k.putRecords(0, records); err != nil { + if err := k.putRecords(ctx, 0, records); err != nil { return ctxerr.Wrap(ctx, err, "put records") } totalBytes = 0 records = nil } - records = append(records, &kinesis.PutRecordsRequestEntry{Data: []byte(log), PartitionKey: aws.String(partitionKey)}) + records = append(records, types.PutRecordsRequestEntry{ + Data: []byte(log), + PartitionKey: aws.String(partitionKey), + }) totalBytes += len(log) + len(partitionKey) } // Push the final batch if len(records) > 0 { - if err := k.putRecords(0, records); err != nil { + if err := k.putRecords(ctx, 0, records); err != nil { return ctxerr.Wrap(ctx, err, "put records") } } @@ -156,7 +172,7 @@ func (k *kinesisLogWriter) Write(ctx context.Context, logs []json.RawMessage) er return nil } -func (k *kinesisLogWriter) putRecords(try int, records []*kinesis.PutRecordsRequestEntry) error { +func (k *kinesisLogWriter) putRecords(ctx context.Context, try int, records []types.PutRecordsRequestEntry) error { if try > 0 { time.Sleep(100 * time.Millisecond * time.Duration(math.Pow(2.0, float64(try)))) } @@ -165,13 +181,13 @@ func (k *kinesisLogWriter) putRecords(try int, records []*kinesis.PutRecordsRequ Records: records, } - output, err := k.client.PutRecords(input) + output, err := k.client.PutRecords(ctx, input) if err != nil { - var ae awserr.Error - if errors.As(err, &ae) { + var anyAPIError smithy.APIError + if errors.As(err, &anyAPIError) { if try < kinesisMaxRetries { // Retry with backoff - return k.putRecords(try+1, records) + return k.putRecords(ctx, try+1, records) } } @@ -199,7 +215,7 @@ func (k *kinesisLogWriter) putRecords(try int, records []*kinesis.PutRecordsRequ ) } - var failedRecords []*kinesis.PutRecordsRequestEntry + var failedRecords []types.PutRecordsRequestEntry // Collect failed records for retry for i, record := range output.Records { if record.ErrorCode != nil { @@ -207,7 +223,7 @@ func (k *kinesisLogWriter) putRecords(try int, records []*kinesis.PutRecordsRequ } } - return k.putRecords(try+1, failedRecords) + return k.putRecords(ctx, try+1, failedRecords) } return nil diff --git a/server/logging/kinesis_test.go b/server/logging/kinesis_test.go index bc7fd1eccf..2292ac0268 100644 --- a/server/logging/kinesis_test.go +++ b/server/logging/kinesis_test.go @@ -9,16 +9,15 @@ import ( "testing" "time" - "github.com/aws/aws-sdk-go/aws" - "github.com/aws/aws-sdk-go/aws/awserr" - "github.com/aws/aws-sdk-go/service/kinesis" - "github.com/aws/aws-sdk-go/service/kinesis/kinesisiface" + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/kinesis" + "github.com/aws/aws-sdk-go-v2/service/kinesis/types" "github.com/fleetdm/fleet/v4/server/logging/mock" "github.com/go-kit/log" "github.com/stretchr/testify/assert" ) -func makeKinesisWriterWithMock(client kinesisiface.KinesisAPI, stream string) *kinesisLogWriter { +func makeKinesisWriterWithMock(client KinesisAPI, stream string) *kinesisLogWriter { return &kinesisLogWriter{ client: client, stream: stream, @@ -39,12 +38,12 @@ func getLogsFromPutRecordsInput(input *kinesis.PutRecordsInput) []json.RawMessag func TestKinesisRetryableFailure(t *testing.T) { ctx := context.Background() callCount := 0 - putFunc := func(input *kinesis.PutRecordsInput) (*kinesis.PutRecordsOutput, error) { + putFunc := func(ctx context.Context, input *kinesis.PutRecordsInput, optFns ...func(*kinesis.Options)) (*kinesis.PutRecordsOutput, error) { callCount += 1 assert.Equal(t, logs, getLogsFromPutRecordsInput(input)) assert.Equal(t, "foobar", *input.StreamName) if callCount < 3 { - return nil, awserr.New(kinesis.ErrCodeProvisionedThroughputExceededException, "", nil) + return nil, &types.ProvisionedThroughputExceededException{} } // Returning a non-retryable error earlier helps keep this test faster return nil, errors.New("generic error") @@ -59,11 +58,11 @@ func TestKinesisRetryableFailure(t *testing.T) { func TestKinesisNormalPut(t *testing.T) { ctx := context.Background() callCount := 0 - putFunc := func(input *kinesis.PutRecordsInput) (*kinesis.PutRecordsOutput, error) { + putFunc := func(ctx context.Context, input *kinesis.PutRecordsInput, optFns ...func(*kinesis.Options)) (*kinesis.PutRecordsOutput, error) { callCount += 1 assert.Equal(t, logs, getLogsFromPutRecordsInput(input)) assert.Equal(t, "foobar", *input.StreamName) - return &kinesis.PutRecordsOutput{FailedRecordCount: aws.Int64(0)}, nil + return &kinesis.PutRecordsOutput{FailedRecordCount: aws.Int32(0)}, nil } k := &mock.KinesisMock{PutRecordsFunc: putFunc} writer := makeKinesisWriterWithMock(k, "foobar") @@ -77,48 +76,48 @@ func TestKinesisSomeFailures(t *testing.T) { k := &mock.KinesisMock{} callCount := 0 - call3 := func(input *kinesis.PutRecordsInput) (*kinesis.PutRecordsOutput, error) { + call3 := func(ctx context.Context, input *kinesis.PutRecordsInput, optFns ...func(*kinesis.Options)) (*kinesis.PutRecordsOutput, error) { // final invocation callCount += 1 assert.Equal(t, logs[1:2], getLogsFromPutRecordsInput(input)) return &kinesis.PutRecordsOutput{ - FailedRecordCount: aws.Int64(0), + FailedRecordCount: aws.Int32(0), }, nil } - call2 := func(input *kinesis.PutRecordsInput) (*kinesis.PutRecordsOutput, error) { + call2 := func(ctx context.Context, input *kinesis.PutRecordsInput, optFns ...func(*kinesis.Options)) (*kinesis.PutRecordsOutput, error) { // Set to invoke call3 next time k.PutRecordsFunc = call3 callCount += 1 assert.Equal(t, logs[1:], getLogsFromPutRecordsInput(input)) return &kinesis.PutRecordsOutput{ - FailedRecordCount: aws.Int64(1), - Records: []*kinesis.PutRecordsResultEntry{ - &kinesis.PutRecordsResultEntry{ + FailedRecordCount: aws.Int32(1), + Records: []types.PutRecordsResultEntry{ + { ErrorCode: aws.String("error"), }, - &kinesis.PutRecordsResultEntry{ + { SequenceNumber: aws.String("foo"), }, }, }, nil } - call1 := func(input *kinesis.PutRecordsInput) (*kinesis.PutRecordsOutput, error) { + call1 := func(ctx context.Context, input *kinesis.PutRecordsInput, optFns ...func(*kinesis.Options)) (*kinesis.PutRecordsOutput, error) { // Use call2 function for next call k.PutRecordsFunc = call2 callCount += 1 assert.Equal(t, logs, getLogsFromPutRecordsInput(input)) return &kinesis.PutRecordsOutput{ - FailedRecordCount: aws.Int64(1), - Records: []*kinesis.PutRecordsResultEntry{ - &kinesis.PutRecordsResultEntry{ + FailedRecordCount: aws.Int32(1), + Records: []types.PutRecordsResultEntry{ + { SequenceNumber: aws.String("foo"), }, - &kinesis.PutRecordsResultEntry{ + { ErrorCode: aws.String("error"), }, - &kinesis.PutRecordsResultEntry{ + { ErrorCode: aws.String("error"), }, }, @@ -136,13 +135,13 @@ func TestKinesisFailAllRecords(t *testing.T) { k := &mock.KinesisMock{} callCount := 0 - k.PutRecordsFunc = func(input *kinesis.PutRecordsInput) (*kinesis.PutRecordsOutput, error) { + k.PutRecordsFunc = func(ctx context.Context, input *kinesis.PutRecordsInput, optFns ...func(*kinesis.Options)) (*kinesis.PutRecordsOutput, error) { callCount += 1 assert.Equal(t, logs, getLogsFromPutRecordsInput(input)) if callCount < 3 { return &kinesis.PutRecordsOutput{ - FailedRecordCount: aws.Int64(1), - Records: []*kinesis.PutRecordsResultEntry{ + FailedRecordCount: aws.Int32(1), + Records: []types.PutRecordsResultEntry{ {ErrorCode: aws.String("error")}, {ErrorCode: aws.String("error")}, {ErrorCode: aws.String("error")}, @@ -166,11 +165,11 @@ func TestKinesisRecordTooBig(t *testing.T) { copy(newLogs, logs) newLogs[0] = make(json.RawMessage, kinesisMaxSizeOfRecord+1) callCount := 0 - putFunc := func(input *kinesis.PutRecordsInput) (*kinesis.PutRecordsOutput, error) { + putFunc := func(ctx context.Context, input *kinesis.PutRecordsInput, optFns ...func(*kinesis.Options)) (*kinesis.PutRecordsOutput, error) { callCount += 1 assert.Equal(t, newLogs[1:], getLogsFromPutRecordsInput(input)) assert.Equal(t, "foobar", *input.StreamName) - return &kinesis.PutRecordsOutput{FailedRecordCount: aws.Int64(0)}, nil + return &kinesis.PutRecordsOutput{FailedRecordCount: aws.Int32(0)}, nil } k := &mock.KinesisMock{PutRecordsFunc: putFunc} writer := makeKinesisWriterWithMock(k, "foobar") @@ -188,11 +187,11 @@ func TestKinesisSplitBatchBySize(t *testing.T) { logs[i] = make(json.RawMessage, kinesisMaxSizeOfRecord-1-256) } callCount := 0 - putFunc := func(input *kinesis.PutRecordsInput) (*kinesis.PutRecordsOutput, error) { + putFunc := func(ctx context.Context, input *kinesis.PutRecordsInput, optFns ...func(*kinesis.Options)) (*kinesis.PutRecordsOutput, error) { callCount += 1 assert.Len(t, getLogsFromPutRecordsInput(input), 5) assert.Equal(t, "foobar", *input.StreamName) - return &kinesis.PutRecordsOutput{FailedRecordCount: aws.Int64(0)}, nil + return &kinesis.PutRecordsOutput{FailedRecordCount: aws.Int32(0)}, nil } k := &mock.KinesisMock{PutRecordsFunc: putFunc} writer := makeKinesisWriterWithMock(k, "foobar") @@ -208,11 +207,11 @@ func TestKinesisSplitBatchByCount(t *testing.T) { logs[i] = json.RawMessage(`{}`) } callCount := 0 - putFunc := func(input *kinesis.PutRecordsInput) (*kinesis.PutRecordsOutput, error) { + putFunc := func(ctx context.Context, input *kinesis.PutRecordsInput, optFns ...func(*kinesis.Options)) (*kinesis.PutRecordsOutput, error) { callCount += 1 assert.Len(t, getLogsFromPutRecordsInput(input), kinesisMaxRecordsInBatch) assert.Equal(t, "foobar", *input.StreamName) - return &kinesis.PutRecordsOutput{FailedRecordCount: aws.Int64(0)}, nil + return &kinesis.PutRecordsOutput{FailedRecordCount: aws.Int32(0)}, nil } k := &mock.KinesisMock{PutRecordsFunc: putFunc} writer := makeKinesisWriterWithMock(k, "foobar") @@ -222,45 +221,45 @@ func TestKinesisSplitBatchByCount(t *testing.T) { } func TestKinesisValidateStreamActive(t *testing.T) { - describeFunc := func(input *kinesis.DescribeStreamInput) (*kinesis.DescribeStreamOutput, error) { + describeFunc := func(ctx context.Context, input *kinesis.DescribeStreamInput, optFns ...func(*kinesis.Options)) (*kinesis.DescribeStreamOutput, error) { assert.Equal(t, "test", *input.StreamName) return &kinesis.DescribeStreamOutput{ - StreamDescription: &kinesis.StreamDescription{ - StreamStatus: aws.String(kinesis.StreamStatusActive), + StreamDescription: &types.StreamDescription{ + StreamStatus: types.StreamStatusActive, }, }, nil } k := &mock.KinesisMock{DescribeStreamFunc: describeFunc} writer := makeKinesisWriterWithMock(k, "test") - err := writer.validateStream() + err := writer.validateStream(context.Background()) assert.NoError(t, err) assert.True(t, k.DescribeStreamFuncInvoked) } func TestKinesisValidateStreamNotActive(t *testing.T) { - describeFunc := func(input *kinesis.DescribeStreamInput) (*kinesis.DescribeStreamOutput, error) { + describeFunc := func(ctx context.Context, input *kinesis.DescribeStreamInput, optFns ...func(*kinesis.Options)) (*kinesis.DescribeStreamOutput, error) { assert.Equal(t, "test", *input.StreamName) return &kinesis.DescribeStreamOutput{ - StreamDescription: &kinesis.StreamDescription{ - StreamStatus: aws.String(kinesis.StreamStatusCreating), + StreamDescription: &types.StreamDescription{ + StreamStatus: types.StreamStatusCreating, }, }, nil } k := &mock.KinesisMock{DescribeStreamFunc: describeFunc} writer := makeKinesisWriterWithMock(k, "test") - err := writer.validateStream() + err := writer.validateStream(context.Background()) assert.Error(t, err) assert.True(t, k.DescribeStreamFuncInvoked) } func TestKinesisValidateStreamError(t *testing.T) { - describeFunc := func(input *kinesis.DescribeStreamInput) (*kinesis.DescribeStreamOutput, error) { + describeFunc := func(ctx context.Context, input *kinesis.DescribeStreamInput, optFns ...func(*kinesis.Options)) (*kinesis.DescribeStreamOutput, error) { assert.Equal(t, "test", *input.StreamName) return nil, errors.New("kaboom!") } k := &mock.KinesisMock{DescribeStreamFunc: describeFunc} writer := makeKinesisWriterWithMock(k, "test") - err := writer.validateStream() + err := writer.validateStream(context.Background()) assert.Error(t, err) assert.True(t, k.DescribeStreamFuncInvoked) } diff --git a/server/logging/lambda.go b/server/logging/lambda.go index 61ec67b836..94bdee8d40 100644 --- a/server/logging/lambda.go +++ b/server/logging/lambda.go @@ -5,12 +5,11 @@ import ( "encoding/json" "fmt" - "github.com/aws/aws-sdk-go/aws" - "github.com/aws/aws-sdk-go/aws/credentials" - "github.com/aws/aws-sdk-go/aws/credentials/stscreds" - "github.com/aws/aws-sdk-go/aws/session" - "github.com/aws/aws-sdk-go/service/lambda" - "github.com/aws/aws-sdk-go/service/lambda/lambdaiface" + aws_config "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/credentials" + "github.com/aws/aws-sdk-go-v2/credentials/stscreds" + "github.com/aws/aws-sdk-go-v2/service/lambda" + "github.com/aws/aws-sdk-go-v2/service/lambda/types" "github.com/go-kit/log" "github.com/go-kit/log/level" ) @@ -24,60 +23,60 @@ const ( lambdaMaxSizeOfPayload = 6 * 1000 * 1000 // 6MB ) +type LambdaAPI interface { + Invoke(ctx context.Context, params *lambda.InvokeInput, optFns ...func(*lambda.Options)) (*lambda.InvokeOutput, error) +} + type lambdaLogWriter struct { - client lambdaiface.LambdaAPI + client LambdaAPI functionName string logger log.Logger } func NewLambdaLogWriter(region, id, secret, stsAssumeRoleArn, stsExternalID, functionName string, logger log.Logger) (*lambdaLogWriter, error) { - conf := &aws.Config{ - Region: ®ion, - } + var opts []func(*aws_config.LoadOptions) error // Only provide static credentials if we have them - // otherwise use the default credentials provider chain + // otherwise use the default credentials provider chain. if id != "" && secret != "" { - conf.Credentials = credentials.NewStaticCredentials(id, secret, "") - } - - sess, err := session.NewSession(conf) - if err != nil { - return nil, fmt.Errorf("create Lambda client: %w", err) + opts = append(opts, + aws_config.WithCredentialsProvider(credentials.NewStaticCredentialsProvider(id, secret, "")), + ) } + // cfg.StsAssumeRoleArn has been marked as deprecated, but we still set it in case users are using it. if stsAssumeRoleArn != "" { - creds := stscreds.NewCredentials(sess, stsAssumeRoleArn, func(provider *stscreds.AssumeRoleProvider) { + opts = append(opts, aws_config.WithAssumeRoleCredentialOptions(func(r *stscreds.AssumeRoleOptions) { + r.RoleARN = stsAssumeRoleArn if stsExternalID != "" { - provider.ExternalID = &stsExternalID + r.ExternalID = &stsExternalID } - }) - conf.Credentials = creds - - sess, err = session.NewSession(conf) - - if err != nil { - return nil, fmt.Errorf("create Lambda client: %w", err) - } + })) } - client := lambda.New(sess) + + opts = append(opts, aws_config.WithRegion(region)) + conf, err := aws_config.LoadDefaultConfig(context.Background(), opts...) + if err != nil { + return nil, fmt.Errorf("failed to create default config: %w", err) + } + lambdaClient := lambda.NewFromConfig(conf) f := &lambdaLogWriter{ - client: client, + client: lambdaClient, functionName: functionName, logger: logger, } - if err := f.validateFunction(); err != nil { + if err := f.validateFunction(context.Background()); err != nil { return nil, fmt.Errorf("validate lambda: %w", err) } return f, nil } -func (f *lambdaLogWriter) validateFunction() error { - out, err := f.client.Invoke( +func (f *lambdaLogWriter) validateFunction(ctx context.Context) error { + out, err := f.client.Invoke(ctx, &lambda.InvokeInput{ FunctionName: &f.functionName, - InvocationType: aws.String("DryRun"), + InvocationType: types.InvocationTypeDryRun, }, ) if err != nil { @@ -108,7 +107,7 @@ func (f *lambdaLogWriter) Write(ctx context.Context, logs []json.RawMessage) err continue } - out, err := f.client.Invoke( + out, err := f.client.Invoke(ctx, &lambda.InvokeInput{ FunctionName: &f.functionName, Payload: []byte(log), diff --git a/server/logging/lambda_test.go b/server/logging/lambda_test.go index 64ab9ff7c2..0c67e6d327 100644 --- a/server/logging/lambda_test.go +++ b/server/logging/lambda_test.go @@ -5,9 +5,9 @@ import ( "errors" "testing" - "github.com/aws/aws-sdk-go/aws" - "github.com/aws/aws-sdk-go/service/lambda" - "github.com/aws/aws-sdk-go/service/lambda/lambdaiface" + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/lambda" + "github.com/aws/aws-sdk-go-v2/service/lambda/types" "github.com/fleetdm/fleet/v4/server/logging/mock" "github.com/fleetdm/fleet/v4/server/test" "github.com/go-kit/log" @@ -15,7 +15,7 @@ import ( tmock "github.com/stretchr/testify/mock" ) -func makeLambdaWriterWithMock(client lambdaiface.LambdaAPI, functionName string) *lambdaLogWriter { +func makeLambdaWriterWithMock(client LambdaAPI, functionName string) *lambdaLogWriter { return &lambdaLogWriter{ client: client, functionName: functionName, @@ -25,30 +25,33 @@ func makeLambdaWriterWithMock(client lambdaiface.LambdaAPI, functionName string) func TestLambdaValidateFunctionError(t *testing.T) { m := &mock.LambdaMock{} - m.On("Invoke", &lambda.InvokeInput{FunctionName: aws.String("foobar"), InvocationType: aws.String("DryRun")}). + ctx := context.Background() + m.On("Invoke", &lambda.InvokeInput{FunctionName: aws.String("foobar"), InvocationType: types.InvocationTypeDryRun}). Return(nil, errors.New("failed")) writer := makeLambdaWriterWithMock(m, "foobar") - err := writer.validateFunction() + err := writer.validateFunction(ctx) assert.Error(t, err) m.AssertExpectations(test.Quiet(t)) } func TestLambdaValidateFunctionErrorFunction(t *testing.T) { m := &mock.LambdaMock{} - m.On("Invoke", &lambda.InvokeInput{FunctionName: aws.String("foobar"), InvocationType: aws.String("DryRun")}). + ctx := context.Background() + m.On("Invoke", &lambda.InvokeInput{FunctionName: aws.String("foobar"), InvocationType: types.InvocationTypeDryRun}). Return(&lambda.InvokeOutput{FunctionError: aws.String("failed")}, nil) writer := makeLambdaWriterWithMock(m, "foobar") - err := writer.validateFunction() + err := writer.validateFunction(ctx) assert.Error(t, err) m.AssertExpectations(test.Quiet(t)) } func TestLambdaValidateFunctionSuccess(t *testing.T) { m := &mock.LambdaMock{} - m.On("Invoke", &lambda.InvokeInput{FunctionName: aws.String("foobar"), InvocationType: aws.String("DryRun")}). + ctx := context.Background() + m.On("Invoke", &lambda.InvokeInput{FunctionName: aws.String("foobar"), InvocationType: types.InvocationTypeDryRun}). Return(&lambda.InvokeOutput{}, nil) writer := makeLambdaWriterWithMock(m, "foobar") - err := writer.validateFunction() + err := writer.validateFunction(ctx) assert.NoError(t, err) m.AssertExpectations(test.Quiet(t)) } @@ -57,7 +60,7 @@ func TestLambdaError(t *testing.T) { m := &mock.LambdaMock{} m.On("Invoke", tmock.MatchedBy( func(in *lambda.InvokeInput) bool { - return *in.FunctionName == "foobar" && in.InvocationType == nil + return *in.FunctionName == "foobar" && in.InvocationType == "" }, )).Return(nil, errors.New("failed")) writer := makeLambdaWriterWithMock(m, "foobar") @@ -70,7 +73,7 @@ func TestLambdaSuccess(t *testing.T) { m := &mock.LambdaMock{} m.On("Invoke", tmock.MatchedBy( func(in *lambda.InvokeInput) bool { - return len(in.Payload) > 0 && *in.FunctionName == "foobar" && in.InvocationType == nil + return len(in.Payload) > 0 && *in.FunctionName == "foobar" && in.InvocationType == "" }, )).Return(&lambda.InvokeOutput{}, nil). Times(len(logs)) diff --git a/server/logging/mock/firehose.go b/server/logging/mock/firehose.go index 8742ecdfa1..92600a4bc9 100644 --- a/server/logging/mock/firehose.go +++ b/server/logging/mock/firehose.go @@ -1,29 +1,29 @@ package mock import ( - "github.com/aws/aws-sdk-go/service/firehose" - "github.com/aws/aws-sdk-go/service/firehose/firehoseiface" + "context" + + "github.com/aws/aws-sdk-go-v2/service/firehose" ) -var _ firehoseiface.FirehoseAPI = (*FirehoseMock)(nil) +type ( + PutRecordBatchFunc func(context.Context, *firehose.PutRecordBatchInput, ...func(*firehose.Options)) (*firehose.PutRecordBatchOutput, error) + DescribeDeliveryStreamFunc func(context.Context, *firehose.DescribeDeliveryStreamInput, ...func(*firehose.Options)) (*firehose.DescribeDeliveryStreamOutput, error) -type PutRecordBatchFunc func(*firehose.PutRecordBatchInput) (*firehose.PutRecordBatchOutput, error) -type DescribeDeliveryStreamFunc func(input *firehose.DescribeDeliveryStreamInput) (*firehose.DescribeDeliveryStreamOutput, error) -type FirehoseMock struct { - firehoseiface.FirehoseAPI + FirehoseMock struct { + PutRecordBatchFunc PutRecordBatchFunc + PutRecordBatchFuncInvoked bool + DescribeDeliveryStreamFunc DescribeDeliveryStreamFunc + DescribeDeliveryStreamFuncInvoked bool + } +) - PutRecordBatchFunc PutRecordBatchFunc - PutRecordBatchFuncInvoked bool - DescribeDeliveryStreamFunc DescribeDeliveryStreamFunc - DescribeDeliveryStreamFuncInvoked bool -} - -func (f *FirehoseMock) PutRecordBatch(input *firehose.PutRecordBatchInput) (*firehose.PutRecordBatchOutput, error) { +func (f *FirehoseMock) PutRecordBatch(ctx context.Context, input *firehose.PutRecordBatchInput, optFns ...func(*firehose.Options)) (*firehose.PutRecordBatchOutput, error) { f.PutRecordBatchFuncInvoked = true - return f.PutRecordBatchFunc(input) + return f.PutRecordBatchFunc(ctx, input, optFns...) } -func (f *FirehoseMock) DescribeDeliveryStream(input *firehose.DescribeDeliveryStreamInput) (*firehose.DescribeDeliveryStreamOutput, error) { +func (f *FirehoseMock) DescribeDeliveryStream(ctx context.Context, input *firehose.DescribeDeliveryStreamInput, optFns ...func(*firehose.Options)) (*firehose.DescribeDeliveryStreamOutput, error) { f.DescribeDeliveryStreamFuncInvoked = true - return f.DescribeDeliveryStreamFunc(input) + return f.DescribeDeliveryStreamFunc(ctx, input, optFns...) } diff --git a/server/logging/mock/kinesis.go b/server/logging/mock/kinesis.go index c7a2b0291a..533fe708a3 100644 --- a/server/logging/mock/kinesis.go +++ b/server/logging/mock/kinesis.go @@ -1,29 +1,28 @@ package mock import ( - "github.com/aws/aws-sdk-go/service/kinesis" - "github.com/aws/aws-sdk-go/service/kinesis/kinesisiface" + "context" + + "github.com/aws/aws-sdk-go-v2/service/kinesis" ) -var _ kinesisiface.KinesisAPI = (*KinesisMock)(nil) +type ( + PutRecordsFunc func(context.Context, *kinesis.PutRecordsInput, ...func(*kinesis.Options)) (*kinesis.PutRecordsOutput, error) + DescribeStreamFunc func(context.Context, *kinesis.DescribeStreamInput, ...func(*kinesis.Options)) (*kinesis.DescribeStreamOutput, error) + KinesisMock struct { + PutRecordsFunc PutRecordsFunc + PutRecordsFuncInvoked bool + DescribeStreamFunc DescribeStreamFunc + DescribeStreamFuncInvoked bool + } +) -type PutRecordsFunc func(*kinesis.PutRecordsInput) (*kinesis.PutRecordsOutput, error) -type DescribeStreamFunc func(input *kinesis.DescribeStreamInput) (*kinesis.DescribeStreamOutput, error) -type KinesisMock struct { - kinesisiface.KinesisAPI - - PutRecordsFunc PutRecordsFunc - PutRecordsFuncInvoked bool - DescribeStreamFunc DescribeStreamFunc - DescribeStreamFuncInvoked bool -} - -func (k *KinesisMock) PutRecords(input *kinesis.PutRecordsInput) (*kinesis.PutRecordsOutput, error) { +func (k *KinesisMock) PutRecords(ctx context.Context, input *kinesis.PutRecordsInput, optFns ...func(*kinesis.Options)) (*kinesis.PutRecordsOutput, error) { k.PutRecordsFuncInvoked = true - return k.PutRecordsFunc(input) + return k.PutRecordsFunc(ctx, input, optFns...) } -func (k *KinesisMock) DescribeStream(input *kinesis.DescribeStreamInput) (*kinesis.DescribeStreamOutput, error) { +func (k *KinesisMock) DescribeStream(ctx context.Context, input *kinesis.DescribeStreamInput, optFns ...func(*kinesis.Options)) (*kinesis.DescribeStreamOutput, error) { k.DescribeStreamFuncInvoked = true - return k.DescribeStreamFunc(input) + return k.DescribeStreamFunc(ctx, input, optFns...) } diff --git a/server/logging/mock/lambda.go b/server/logging/mock/lambda.go index 16a6fb4721..b7eedd88f7 100644 --- a/server/logging/mock/lambda.go +++ b/server/logging/mock/lambda.go @@ -1,17 +1,17 @@ package mock import ( - "github.com/aws/aws-sdk-go/service/lambda" - "github.com/aws/aws-sdk-go/service/lambda/lambdaiface" + "context" + + "github.com/aws/aws-sdk-go-v2/service/lambda" "github.com/stretchr/testify/mock" ) type LambdaMock struct { mock.Mock - lambdaiface.LambdaAPI } -func (l *LambdaMock) Invoke(input *lambda.InvokeInput) (*lambda.InvokeOutput, error) { +func (l *LambdaMock) Invoke(ctx context.Context, input *lambda.InvokeInput, optFns ...func(*lambda.Options)) (*lambda.InvokeOutput, error) { args := l.Called(input) out, err := args.Get(0), args.Error(1) if out == nil { diff --git a/server/mail/mail.go b/server/mail/mail.go index 277ab5b942..7dc9fcd1ee 100644 --- a/server/mail/mail.go +++ b/server/mail/mail.go @@ -3,6 +3,7 @@ package mail import ( "bytes" + "context" "crypto/tls" "errors" "fmt" @@ -37,7 +38,7 @@ func NewService(config config.FleetConfig) (fleet.MailService, error) { type mailService struct{} type sender interface { - sendMail(e fleet.Email, msg []byte) error + sendMail(ctx context.Context, e fleet.Email, msg []byte) error } func Test(mailer fleet.MailService, e fleet.Email) error { @@ -51,7 +52,7 @@ func Test(mailer fleet.MailService, e fleet.Email) error { return nil } - err = svc.sendMail(e, mailBody) + err = svc.sendMail(context.Background(), e, mailBody) if err != nil { return fmt.Errorf("sending mail: %w", err) } @@ -92,7 +93,7 @@ func getFrom(e fleet.Email) (string, error) { return "From: " + e.SMTPSettings.SMTPSenderAddress + "\r\n", nil } -func (m mailService) SendEmail(e fleet.Email) error { +func (m mailService) SendEmail(ctx context.Context, e fleet.Email) error { if !e.SMTPSettings.SMTPConfigured { return errors.New("email not configured") } @@ -100,7 +101,7 @@ func (m mailService) SendEmail(e fleet.Email) error { if err != nil { return err } - return m.sendMail(e, msg) + return m.sendMail(ctx, e, msg) } func (m mailService) CanSendEmail(smtpSettings fleet.SMTPSettings) bool { @@ -173,7 +174,7 @@ func smtpAuth(e fleet.Email) (smtp.Auth, error) { return auth, nil } -func (m mailService) sendMail(e fleet.Email, msg []byte) error { +func (m mailService) sendMail(ctx context.Context, e fleet.Email, msg []byte) error { smtpHost := fmt.Sprintf( "%s:%d", e.SMTPSettings.SMTPServer, e.SMTPSettings.SMTPPort) auth, err := smtpAuth(e) diff --git a/server/mail/mail_test.go b/server/mail/mail_test.go index 271d157c08..110072f9f5 100644 --- a/server/mail/mail_test.go +++ b/server/mail/mail_test.go @@ -1,6 +1,7 @@ package mail import ( + "context" "encoding/json" "fmt" "io" @@ -86,7 +87,7 @@ func testSMTPPlainAuth(t *testing.T, mailer fleet.MailService) { }, } - err := mailer.SendEmail(mail) + err := mailer.SendEmail(context.Background(), mail) assert.Nil(t, err) } @@ -112,7 +113,7 @@ func testSMTPPlainAuthInvalidCreds(t *testing.T, mailer fleet.MailService) { }, } - err := mailer.SendEmail(mail) + err := mailer.SendEmail(context.Background(), mail) assert.Error(t, err) } @@ -138,7 +139,7 @@ func testSMTPSkipVerify(t *testing.T, mailer fleet.MailService) { }, } - err := mailer.SendEmail(mail) + err := mailer.SendEmail(context.Background(), mail) assert.Nil(t, err) } @@ -161,7 +162,7 @@ func testSMTPNoAuthWithTLS(t *testing.T, mailer fleet.MailService) { }, } - err := mailer.SendEmail(mail) + err := mailer.SendEmail(context.Background(), mail) assert.Nil(t, err) } @@ -190,7 +191,7 @@ func testSMTPDomain(t *testing.T, mailer fleet.MailService) { }, } - err := mailer.SendEmail(mail) + err := mailer.SendEmail(context.Background(), mail) assert.Nil(t, err) rawMsg := getLastRawMailpitMessageFrom(t, randomAddress) diff --git a/server/mail/ses.go b/server/mail/ses.go index cb0eb442bf..91e89c02bc 100644 --- a/server/mail/ses.go +++ b/server/mail/ses.go @@ -1,20 +1,22 @@ package mail import ( + "context" "errors" "fmt" "net/url" - "github.com/aws/aws-sdk-go/aws" - "github.com/aws/aws-sdk-go/aws/credentials" - "github.com/aws/aws-sdk-go/aws/credentials/stscreds" - "github.com/aws/aws-sdk-go/aws/session" - "github.com/aws/aws-sdk-go/service/ses" + "github.com/aws/aws-sdk-go-v2/aws" + aws_config "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/credentials" + "github.com/aws/aws-sdk-go-v2/credentials/stscreds" + "github.com/aws/aws-sdk-go-v2/service/ses" + "github.com/aws/aws-sdk-go-v2/service/ses/types" "github.com/fleetdm/fleet/v4/server/fleet" ) type fleetSESSender interface { - SendRawEmail(input *ses.SendRawEmailInput) (*ses.SendRawEmailOutput, error) + SendRawEmail(ctx context.Context, input *ses.SendRawEmailInput, optFns ...func(*ses.Options)) (*ses.SendRawEmailOutput, error) } type sesSender struct { @@ -30,7 +32,7 @@ func getFromSES(e fleet.Email) (string, error) { return fmt.Sprintf("From: %s\r\n", fmt.Sprintf("do-not-reply@%s", serverURL.Host)), nil } -func (s *sesSender) SendEmail(e fleet.Email) error { +func (s *sesSender) SendEmail(ctx context.Context, e fleet.Email) error { if s.client == nil { return errors.New("ses sender not configured") } @@ -38,7 +40,7 @@ func (s *sesSender) SendEmail(e fleet.Email) error { if err != nil { return err } - return s.sendMail(e, msg) + return s.sendMail(ctx, e, msg) } func (s *sesSender) CanSendEmail(smtpSettings fleet.SMTPSettings) bool { @@ -46,50 +48,57 @@ func (s *sesSender) CanSendEmail(smtpSettings fleet.SMTPSettings) bool { } func NewSESSender(region, endpointURL, id, secret, stsAssumeRoleArn, stsExternalID, sourceArn string) (*sesSender, error) { - conf := &aws.Config{ - Region: ®ion, - Endpoint: &endpointURL, // empty string or nil will use default values + var opts []func(*aws_config.LoadOptions) error + + // The service endpoint is deprecated, but we still set it + // in case users are using it. + if endpointURL != "" { + opts = append(opts, aws_config.WithEndpointResolver(aws.EndpointResolverFunc( + func(service, region string) (aws.Endpoint, error) { + return aws.Endpoint{ + URL: endpointURL, + }, nil + })), + ) } // Only provide static credentials if we have them - // otherwise use the default credentials provider chain + // otherwise use the default credentials provider chain. if id != "" && secret != "" { - conf.Credentials = credentials.NewStaticCredentials(id, secret, "") - } - - sess, err := session.NewSession(conf) - if err != nil { - return nil, fmt.Errorf("create SES client: %w", err) + opts = append(opts, + aws_config.WithCredentialsProvider(credentials.NewStaticCredentialsProvider(id, secret, "")), + ) } + // cfg.StsAssumeRoleArn has been marked as deprecated, but we still set it in case users are using it. if stsAssumeRoleArn != "" { - creds := stscreds.NewCredentials(sess, stsAssumeRoleArn, func(provider *stscreds.AssumeRoleProvider) { + opts = append(opts, aws_config.WithAssumeRoleCredentialOptions(func(r *stscreds.AssumeRoleOptions) { + r.RoleARN = stsAssumeRoleArn if stsExternalID != "" { - provider.ExternalID = &stsExternalID + r.ExternalID = &stsExternalID } - }) - conf.Credentials = creds - - sess, err = session.NewSession(conf) - - if err != nil { - return nil, fmt.Errorf("create SES client: %w", err) - } + })) } - return &sesSender{client: ses.New(sess), sourceArn: sourceArn}, nil + + opts = append(opts, aws_config.WithRegion(region)) + conf, err := aws_config.LoadDefaultConfig(context.Background(), opts...) + if err != nil { + return nil, fmt.Errorf("failed to create default config: %w", err) + } + + sesClient := ses.NewFromConfig(conf) + + return &sesSender{ + client: sesClient, + sourceArn: sourceArn, + }, nil } -func (s *sesSender) sendMail(e fleet.Email, msg []byte) error { - toAddresses := make([]*string, len(e.To)) - for i := range e.To { - t := e.To[i] - toAddresses[i] = &t - } - - _, err := s.client.SendRawEmail(&ses.SendRawEmailInput{ - Destinations: toAddresses, +func (s *sesSender) sendMail(ctx context.Context, e fleet.Email, msg []byte) error { + _, err := s.client.SendRawEmail(ctx, &ses.SendRawEmailInput{ + Destinations: e.To, FromArn: &s.sourceArn, - RawMessage: &ses.RawMessage{Data: msg}, + RawMessage: &types.RawMessage{Data: msg}, SourceArn: &s.sourceArn, }) if err != nil { diff --git a/server/mail/ses_test.go b/server/mail/ses_test.go index 052c1705d5..96a8440177 100644 --- a/server/mail/ses_test.go +++ b/server/mail/ses_test.go @@ -1,11 +1,12 @@ package mail import ( + "context" "errors" "fmt" "testing" - "github.com/aws/aws-sdk-go/service/ses" + "github.com/aws/aws-sdk-go-v2/service/ses" "github.com/fleetdm/fleet/v4/server/fleet" "github.com/stretchr/testify/assert" ) @@ -52,7 +53,7 @@ type mockSESSender struct { shouldErr bool } -func (m mockSESSender) SendRawEmail(input *ses.SendRawEmailInput) (*ses.SendRawEmailOutput, error) { +func (m mockSESSender) SendRawEmail(ctx context.Context, input *ses.SendRawEmailInput, optFns ...func(*ses.Options)) (*ses.SendRawEmailOutput, error) { if m.shouldErr { return nil, errors.New("some error") } @@ -143,7 +144,7 @@ func Test_sesSender_SendEmail(t *testing.T) { client: tt.fields.client, sourceArn: tt.fields.sourceArn, } - tt.wantErr(t, s.SendEmail(tt.args.e), fmt.Sprintf("SendEmail(%v)", tt.args.e)) + tt.wantErr(t, s.SendEmail(context.Background(), tt.args.e), fmt.Sprintf("SendEmail(%v)", tt.args.e)) }) } } diff --git a/server/service/appconfig.go b/server/service/appconfig.go index 58aad60d18..557cee02ba 100644 --- a/server/service/appconfig.go +++ b/server/service/appconfig.go @@ -2010,10 +2010,6 @@ func (svc *Service) ApplyEnrollSecretSpec(ctx context.Context, spec *fleet.Enrol } } - if svc.config.Packaging.GlobalEnrollSecret != "" { - return ctxerr.New(ctx, "enroll secret cannot be changed when fleet_packaging.global_enroll_secret is set") - } - if applyOpts.DryRun { for _, s := range spec.Secrets { available, err := svc.ds.IsEnrollSecretAvailable(ctx, s.Secret, false, nil) diff --git a/server/service/appconfig_test.go b/server/service/appconfig_test.go index 536dac521a..25e883c30b 100644 --- a/server/service/appconfig_test.go +++ b/server/service/appconfig_test.go @@ -349,15 +349,6 @@ func TestApplyEnrollSecretWithGlobalEnrollConfig(t *testing.T) { ) require.True(t, ds.ApplyEnrollSecretsFuncInvoked) require.NoError(t, err) - - // try to change the enroll secret with the config set - ds.ApplyEnrollSecretsFuncInvoked = false - cfg.Packaging.GlobalEnrollSecret = "xyz" - svc, ctx = newTestServiceWithConfig(t, ds, cfg, nil, nil) - ctx = test.UserContext(ctx, test.UserAdmin) - err = svc.ApplyEnrollSecretSpec(ctx, &fleet.EnrollSecretSpec{Secrets: []*fleet.EnrollSecret{{Secret: "DEF"}}}, fleet.ApplySpecOptions{}) - require.Error(t, err) - require.False(t, ds.ApplyEnrollSecretsFuncInvoked) } func TestCertificateChain(t *testing.T) { diff --git a/server/service/invites.go b/server/service/invites.go index 9f5c1ff763..b60f9cb550 100644 --- a/server/service/invites.go +++ b/server/service/invites.go @@ -129,7 +129,7 @@ func (svc *Service) InviteNewUser(ctx context.Context, payload fleet.InvitePaylo }, } - err = svc.mailService.SendEmail(inviteEmail) + err = svc.mailService.SendEmail(ctx, inviteEmail) if err != nil { return nil, err } diff --git a/server/service/service.go b/server/service/service.go index b3cfa86e41..5567e1be87 100644 --- a/server/service/service.go +++ b/server/service/service.go @@ -179,8 +179,8 @@ func NewService( return validationMiddleware{svc, ds, sso}, nil } -func (svc *Service) SendEmail(mail fleet.Email) error { - return svc.mailService.SendEmail(mail) +func (svc *Service) SendEmail(ctx context.Context, mail fleet.Email) error { + return svc.mailService.SendEmail(ctx, mail) } type validationMiddleware struct { diff --git a/server/service/service_appconfig.go b/server/service/service_appconfig.go index ae61c91afd..a9417affc2 100644 --- a/server/service/service_appconfig.go +++ b/server/service/service_appconfig.go @@ -25,12 +25,9 @@ func (svc *Service) NewAppConfig(ctx context.Context, p fleet.AppConfig) (*fleet } // Set up a default enroll secret - secret := svc.config.Packaging.GlobalEnrollSecret - if secret == "" { - secret, err = server.GenerateRandomText(fleet.EnrollSecretDefaultLength) - if err != nil { - return nil, ctxerr.Wrap(ctx, err, "generate enroll secret string") - } + secret, err := server.GenerateRandomText(fleet.EnrollSecretDefaultLength) + if err != nil { + return nil, ctxerr.Wrap(ctx, err, "generate enroll secret string") } secrets := []*fleet.EnrollSecret{ { diff --git a/server/service/service_appconfig_test.go b/server/service/service_appconfig_test.go index d83479e5f1..81a7d733f3 100644 --- a/server/service/service_appconfig_test.go +++ b/server/service/service_appconfig_test.go @@ -125,30 +125,6 @@ func TestEmptyEnrollSecret(t *testing.T) { require.NoError(t, err) } -func TestNewAppConfigWithGlobalEnrollConfig(t *testing.T) { - ds := new(mock.Store) - cfg := config.TestConfig() - cfg.Packaging.GlobalEnrollSecret = "xyz" - svc, ctx := newTestServiceWithConfig(t, ds, cfg, nil, nil) - - ds.NewAppConfigFunc = func(ctx context.Context, config *fleet.AppConfig) (*fleet.AppConfig, error) { - return config, nil - } - - var gotSecrets []*fleet.EnrollSecret - ds.ApplyEnrollSecretsFunc = func(ctx context.Context, teamID *uint, secrets []*fleet.EnrollSecret) error { - gotSecrets = secrets - return nil - } - - ctx = test.UserContext(ctx, test.UserAdmin) - _, err := svc.NewAppConfig(ctx, fleet.AppConfig{ServerSettings: fleet.ServerSettings{ServerURL: "https://acme.co"}}) - require.NoError(t, err) - require.NotNil(t, gotSecrets) - require.Len(t, gotSecrets, 1) - require.Equal(t, gotSecrets[0].Secret, "xyz") -} - func TestService_LoggingConfig(t *testing.T) { logFile := "/dev/null" if runtime.GOOS == "windows" { diff --git a/server/service/sessions.go b/server/service/sessions.go index 25941160ef..d032b1440d 100644 --- a/server/service/sessions.go +++ b/server/service/sessions.go @@ -162,10 +162,13 @@ func loginEndpoint(ctx context.Context, request interface{}, svc fleet.Service) //goland:noinspection GoErrorStringFormat var sendingMFAEmail = errors.New("sending MFA email") -var noMFASupported = errors.New("client with no MFA email support") -var mfaNotSupportedForClient = endpoint_utils.BadRequestErr( - "Your login client does not support MFA. Please log in via the web, then use an API token to authenticate.", - noMFASupported, + +var ( + noMFASupported = errors.New("client with no MFA email support") + mfaNotSupportedForClient = endpoint_utils.BadRequestErr( + "Your login client does not support MFA. Please log in via the web, then use an API token to authenticate.", + noMFASupported, + ) ) func (svc *Service) Login(ctx context.Context, email, password string, supportsEmailVerification bool) (*fleet.User, *fleet.Session, error) { @@ -705,7 +708,7 @@ func (svc *Service) makeMFAEmail(ctx context.Context, user fleet.User) error { }, } - return svc.mailService.SendEmail(email) + return svc.mailService.SendEmail(ctx, email) } func (svc *Service) GetSessionByKey(ctx context.Context, key string) (*fleet.Session, error) { diff --git a/server/service/testing_utils.go b/server/service/testing_utils.go index cfe5dd9499..ccc1976353 100644 --- a/server/service/testing_utils.go +++ b/server/service/testing_utils.go @@ -318,7 +318,7 @@ type mockMailService struct { Invoked bool } -func (svc *mockMailService) SendEmail(e fleet.Email) error { +func (svc *mockMailService) SendEmail(ctx context.Context, e fleet.Email) error { svc.Invoked = true return svc.SendEmailFn(e) } diff --git a/server/service/users.go b/server/service/users.go index 6b84d6a41d..a6103eff67 100644 --- a/server/service/users.go +++ b/server/service/users.go @@ -937,7 +937,7 @@ func (svc *Service) modifyEmailAddress(ctx context.Context, user *fleet.User, em AssetURL: getAssetURL(), }, } - return svc.mailService.SendEmail(changeEmail) + return svc.mailService.SendEmail(ctx, changeEmail) } // saves user in datastore. @@ -1202,7 +1202,7 @@ func (svc *Service) RequestPasswordReset(ctx context.Context, email string) erro }, } - err = svc.mailService.SendEmail(resetEmail) + err = svc.mailService.SendEmail(ctx, resetEmail) if err != nil { level.Error(svc.logger).Log("err", err, "msg", "failed to send password reset request email") }