diff --git a/pkg/providers/aws/aws.go b/pkg/providers/aws/aws.go
index 61e1dafe..762f5a43 100644
--- a/pkg/providers/aws/aws.go
+++ b/pkg/providers/aws/aws.go
@@ -20,6 +20,7 @@ import (
"github.com/aws/aws-sdk-go/service/lambda"
"github.com/aws/aws-sdk-go/service/lightsail"
"github.com/aws/aws-sdk-go/service/organizations"
+ "github.com/aws/aws-sdk-go/service/rds"
"github.com/aws/aws-sdk-go/service/route53"
"github.com/aws/aws-sdk-go/service/s3"
"github.com/aws/aws-sdk-go/service/sts"
@@ -29,7 +30,7 @@ import (
sliceutil "github.com/projectdiscovery/utils/slice"
)
-var Services = []string{"ec2", "instance", "route53", "s3", "ecs", "eks", "lambda", "apigateway", "apigatewayv2", "alb", "elb", "lightsail", "cloudfront"}
+var Services = []string{"ec2", "instance", "route53", "s3", "ecs", "eks", "lambda", "apigateway", "apigatewayv2", "alb", "elb", "lightsail", "cloudfront", "rds"}
type ProviderOptions struct {
Id string
@@ -113,6 +114,7 @@ type Provider struct {
elbClient *elb.ELB
lightsailClient *lightsail.Lightsail
cloudFrontClient *cloudfront.CloudFront
+ rdsClient *rds.RDS
regions *ec2.DescribeRegionsOutput
session *session.Session
}
@@ -391,6 +393,9 @@ func (p *Provider) initServices(sess *session.Session) {
if services.Has("cloudfront") {
p.cloudFrontClient = cloudfront.New(sess)
}
+ if services.Has("rds") {
+ p.rdsClient = rds.New(sess)
+ }
}
const providerName = "aws"
@@ -494,6 +499,10 @@ func (p *Provider) Resources(ctx context.Context) (*schema.Resources, error) {
cloudfrontProvider := &cloudfrontProvider{cloudFrontClient: p.cloudFrontClient, options: *p.options, session: p.session}
assignWorker(cloudfrontProvider.GetResource)
}
+ if p.rdsClient != nil {
+ rdsProvider := &rdsProvider{rdsClient: p.rdsClient, options: *p.options, session: p.session, regions: p.regions}
+ assignWorker(rdsProvider.GetResource)
+ }
go func() {
workersWaitGroup.Wait()
@@ -648,6 +657,13 @@ func (p *Provider) verify() error {
}
}
+ if !success && p.rdsClient != nil {
+ _, err := p.rdsClient.DescribeDBInstances(&rds.DescribeDBInstancesInput{MaxRecords: aws.Int64(20)})
+ if err == nil {
+ success = true
+ }
+ }
+
if success {
return nil
}
diff --git a/pkg/providers/aws/rds.go b/pkg/providers/aws/rds.go
new file mode 100644
index 00000000..11e2217e
--- /dev/null
+++ b/pkg/providers/aws/rds.go
@@ -0,0 +1,235 @@
+package aws
+
+import (
+ "context"
+ stderrors "errors"
+ "fmt"
+ "strconv"
+ "strings"
+ "sync"
+ "time"
+
+ "github.com/aws/aws-sdk-go/aws"
+ "github.com/aws/aws-sdk-go/aws/credentials/stscreds"
+ "github.com/aws/aws-sdk-go/aws/session"
+ "github.com/aws/aws-sdk-go/service/ec2"
+ "github.com/aws/aws-sdk-go/service/rds"
+ "github.com/pkg/errors"
+ "github.com/projectdiscovery/cloudlist/pkg/schema"
+ "github.com/projectdiscovery/gologger"
+)
+
+// rdsProvider is a provider for AWS RDS API
+type rdsProvider struct {
+ options ProviderOptions
+ rdsClient *rds.RDS
+ session *session.Session
+ regions *ec2.DescribeRegionsOutput
+}
+
+func (rp *rdsProvider) name() string {
+ return "rds"
+}
+
+// GetResource returns all the resources in the store for a provider.
+func (rp *rdsProvider) GetResource(ctx context.Context) (*schema.Resources, error) {
+ list := schema.NewResources()
+ var wg sync.WaitGroup
+ var mu sync.Mutex
+ var errs []error
+
+ for _, region := range rp.regions.Regions {
+ for _, rdsClient := range rp.getRdsClients(region.RegionName) {
+ wg.Add(1)
+
+ go func(client *rds.RDS) {
+ defer wg.Done()
+ defer func() {
+ if r := recover(); r != nil {
+ mu.Lock()
+ errs = append(errs, fmt.Errorf("panic in rds provider: %v", r))
+ mu.Unlock()
+ }
+ }()
+
+ resources, err := rp.listRDSResources(ctx, client)
+ mu.Lock()
+ defer mu.Unlock()
+ if resources != nil {
+ list.Merge(resources)
+ }
+ if err != nil {
+ errs = append(errs, err)
+ }
+ }(rdsClient)
+ }
+ }
+ wg.Wait()
+ if len(errs) > 0 && len(list.Items) == 0 {
+ return nil, fmt.Errorf("rds: all workers failed: %v", errs)
+ }
+ if len(errs) > 0 {
+ gologger.Warning().Msgf("rds: some listings failed: %v", errs)
+ }
+ return list, nil
+}
+
+func (rp *rdsProvider) listRDSResources(ctx context.Context, rdsClient *rds.RDS) (*schema.Resources, error) {
+ list := schema.NewResources()
+ var errs []error
+
+ err := rdsClient.DescribeDBInstancesPagesWithContext(ctx, &rds.DescribeDBInstancesInput{}, func(page *rds.DescribeDBInstancesOutput, _ bool) bool {
+ for _, instance := range page.DBInstances {
+ if instance == nil || instance.Endpoint == nil || aws.StringValue(instance.Endpoint.Address) == "" {
+ continue
+ }
+
+ var metadata map[string]string
+ if rp.options.ExtendedMetadata {
+ metadata = getDBInstanceMetadata(instance)
+ }
+
+ list.Append(&schema.Resource{
+ ID: rp.options.Id,
+ Provider: providerName,
+ DNSName: aws.StringValue(instance.Endpoint.Address),
+ Public: true,
+ Service: rp.name(),
+ Metadata: metadata,
+ })
+ }
+ return true
+ })
+ if err != nil {
+ errs = append(errs, errors.Wrap(err, "could not describe RDS instances"))
+ }
+
+ err = rdsClient.DescribeDBClustersPagesWithContext(ctx, &rds.DescribeDBClustersInput{}, func(page *rds.DescribeDBClustersOutput, _ bool) bool {
+ for _, cluster := range page.DBClusters {
+ if cluster == nil {
+ continue
+ }
+
+ var metadata map[string]string
+ if rp.options.ExtendedMetadata {
+ metadata = getDBClusterMetadata(cluster)
+ }
+
+ endpoints := append([]*string{cluster.Endpoint, cluster.ReaderEndpoint}, cluster.CustomEndpoints...)
+ for _, endpoint := range endpoints {
+ if aws.StringValue(endpoint) == "" {
+ continue
+ }
+ list.Append(&schema.Resource{
+ ID: rp.options.Id,
+ Provider: providerName,
+ DNSName: aws.StringValue(endpoint),
+ Public: true,
+ Service: rp.name(),
+ Metadata: metadata,
+ })
+ }
+ }
+ return true
+ })
+ if err != nil {
+ errs = append(errs, errors.Wrap(err, "could not describe RDS clusters"))
+ }
+ if len(errs) > 0 {
+ return list, stderrors.Join(errs...)
+ }
+ return list, nil
+}
+
+func getDBInstanceMetadata(instance *rds.DBInstance) map[string]string {
+ metadata := make(map[string]string)
+
+ schema.AddMetadata(metadata, "db_instance_identifier", instance.DBInstanceIdentifier)
+ schema.AddMetadata(metadata, "db_instance_arn", instance.DBInstanceArn)
+ schema.AddMetadata(metadata, "db_instance_class", instance.DBInstanceClass)
+ schema.AddMetadata(metadata, "db_instance_status", instance.DBInstanceStatus)
+ schema.AddMetadata(metadata, "db_cluster_identifier", instance.DBClusterIdentifier)
+ schema.AddMetadata(metadata, "engine", instance.Engine)
+ schema.AddMetadata(metadata, "engine_version", instance.EngineVersion)
+ schema.AddMetadata(metadata, "availability_zone", instance.AvailabilityZone)
+ if instance.Endpoint != nil && instance.Endpoint.Port != nil {
+ metadata["port"] = strconv.FormatInt(aws.Int64Value(instance.Endpoint.Port), 10)
+ }
+ if instance.DBSubnetGroup != nil {
+ schema.AddMetadata(metadata, "vpc_id", instance.DBSubnetGroup.VpcId)
+ }
+ // RDS DNS names are always emitted; this tells reachable endpoints apart from VPC-only ones.
+ metadata["publicly_accessible"] = strconv.FormatBool(aws.BoolValue(instance.PubliclyAccessible))
+ if instance.InstanceCreateTime != nil {
+ metadata["create_time"] = instance.InstanceCreateTime.Format(time.RFC3339)
+ }
+ if tagString := buildRDSTagString(instance.TagList); tagString != "" {
+ metadata["tags"] = tagString
+ }
+ return metadata
+}
+
+func getDBClusterMetadata(cluster *rds.DBCluster) map[string]string {
+ metadata := make(map[string]string)
+
+ schema.AddMetadata(metadata, "db_cluster_identifier", cluster.DBClusterIdentifier)
+ schema.AddMetadata(metadata, "db_cluster_arn", cluster.DBClusterArn)
+ schema.AddMetadata(metadata, "db_cluster_status", cluster.Status)
+ schema.AddMetadata(metadata, "engine", cluster.Engine)
+ schema.AddMetadata(metadata, "engine_version", cluster.EngineVersion)
+ schema.AddMetadata(metadata, "engine_mode", cluster.EngineMode)
+ if cluster.Port != nil {
+ metadata["port"] = strconv.FormatInt(aws.Int64Value(cluster.Port), 10)
+ }
+ // Only set for Multi-AZ DB clusters; Aurora tracks it per member instance.
+ if cluster.PubliclyAccessible != nil {
+ metadata["publicly_accessible"] = strconv.FormatBool(aws.BoolValue(cluster.PubliclyAccessible))
+ }
+ if cluster.ClusterCreateTime != nil {
+ metadata["create_time"] = cluster.ClusterCreateTime.Format(time.RFC3339)
+ }
+ if tagString := buildRDSTagString(cluster.TagList); tagString != "" {
+ metadata["tags"] = tagString
+ }
+ return metadata
+}
+
+func buildRDSTagString(tags []*rds.Tag) string {
+ var tagPairs []string
+ for _, tag := range tags {
+ if tag != nil && tag.Key != nil && tag.Value != nil {
+ tagPairs = append(tagPairs, fmt.Sprintf("%s=%s", aws.StringValue(tag.Key), aws.StringValue(tag.Value)))
+ }
+ }
+ return strings.Join(tagPairs, ",")
+}
+
+func (rp *rdsProvider) getRdsClients(region *string) []*rds.RDS {
+ rdsClients := make([]*rds.RDS, 0)
+
+ rdsClient := rds.New(
+ rp.session,
+ aws.NewConfig().WithRegion(aws.StringValue(region)),
+ )
+ rdsClients = append(rdsClients, rdsClient)
+
+ if rp.options.AssumeRoleName == "" || len(rp.options.AccountIds) < 1 {
+ return rdsClients
+ }
+
+ for _, accountId := range rp.options.AccountIds {
+ roleARN := fmt.Sprintf("arn:aws:iam::%s:role/%s", accountId, rp.options.AssumeRoleName)
+ creds := stscreds.NewCredentials(rp.session, roleARN)
+
+ assumeSession, err := session.NewSession(&aws.Config{
+ Region: region,
+ Credentials: creds,
+ })
+ if err != nil {
+ continue
+ }
+
+ rdsClients = append(rdsClients, rds.New(assumeSession))
+ }
+ return rdsClients
+}
diff --git a/pkg/providers/aws/rds_test.go b/pkg/providers/aws/rds_test.go
new file mode 100644
index 00000000..7c681e7a
--- /dev/null
+++ b/pkg/providers/aws/rds_test.go
@@ -0,0 +1,222 @@
+package aws
+
+import (
+ "context"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+
+ "github.com/aws/aws-sdk-go/aws"
+ "github.com/aws/aws-sdk-go/aws/credentials"
+ "github.com/aws/aws-sdk-go/aws/session"
+ "github.com/aws/aws-sdk-go/service/ec2"
+ "github.com/aws/aws-sdk-go/service/rds"
+ "github.com/stretchr/testify/assert"
+ "github.com/stretchr/testify/require"
+)
+
+const describeDBInstancesPage1 = `
+
+ page2
+
+
+ public-db
+ postgres
+ true
+
+ public-db.abc123.us-east-1.rds.amazonaws.com
+ 5432
+
+ envprod
+
+
+ creating-db
+
+
+
+`
+
+const describeDBInstancesPage2 = `
+
+
+
+ private-db
+ mysql
+ false
+
+ private-db.abc123.us-east-1.rds.amazonaws.com
+ 3306
+
+
+
+
+`
+
+const describeDBClustersResponse = `
+
+
+
+ aurora-cluster
+ aurora-postgresql
+ 5432
+ aurora-cluster.cluster-abc123.us-east-1.rds.amazonaws.com
+ aurora-cluster.cluster-ro-abc123.us-east-1.rds.amazonaws.com
+
+ analytics.cluster-custom-abc123.us-east-1.rds.amazonaws.com
+
+
+
+
+`
+
+func newTestRDSClient(t *testing.T) *rds.RDS {
+ t.Helper()
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ require.NoError(t, r.ParseForm())
+ w.Header().Set("Content-Type", "text/xml")
+ switch r.Form.Get("Action") {
+ case "DescribeDBInstances":
+ if r.Form.Get("Marker") == "page2" {
+ _, _ = w.Write([]byte(describeDBInstancesPage2))
+ return
+ }
+ _, _ = w.Write([]byte(describeDBInstancesPage1))
+ case "DescribeDBClusters":
+ _, _ = w.Write([]byte(describeDBClustersResponse))
+ default:
+ t.Errorf("unexpected action %q", r.Form.Get("Action"))
+ w.WriteHeader(http.StatusBadRequest)
+ }
+ }))
+ t.Cleanup(server.Close)
+
+ sess, err := session.NewSession(&aws.Config{
+ Region: aws.String("us-east-1"),
+ Endpoint: aws.String(server.URL),
+ Credentials: credentials.NewStaticCredentials("test", "test", ""),
+ })
+ require.NoError(t, err)
+ return rds.New(sess)
+}
+
+func TestListRDSResources(t *testing.T) {
+ provider := &rdsProvider{options: ProviderOptions{Id: "test", ExtendedMetadata: true}}
+
+ resources, err := provider.listRDSResources(context.Background(), newTestRDSClient(t))
+ require.NoError(t, err)
+
+ byDNS := make(map[string]map[string]string)
+ for _, item := range resources.Items {
+ assert.Equal(t, "rds", item.Service)
+ byDNS[item.DNSName] = item.Metadata
+ }
+ require.Len(t, byDNS, 5, "instances without an endpoint must be skipped, pagination must be followed")
+
+ public := byDNS["public-db.abc123.us-east-1.rds.amazonaws.com"]
+ require.NotNil(t, public)
+ assert.Equal(t, "true", public["publicly_accessible"])
+ assert.Equal(t, "5432", public["port"])
+ assert.Equal(t, "env=prod", public["tags"])
+
+ private := byDNS["private-db.abc123.us-east-1.rds.amazonaws.com"]
+ require.NotNil(t, private)
+ assert.Equal(t, "false", private["publicly_accessible"])
+
+ for _, endpoint := range []string{
+ "aurora-cluster.cluster-abc123.us-east-1.rds.amazonaws.com",
+ "aurora-cluster.cluster-ro-abc123.us-east-1.rds.amazonaws.com",
+ "analytics.cluster-custom-abc123.us-east-1.rds.amazonaws.com",
+ } {
+ require.Contains(t, byDNS, endpoint)
+ assert.Equal(t, "aurora-cluster", byDNS[endpoint]["db_cluster_identifier"])
+ }
+}
+
+func TestRDSGetResourceReportsListingErrors(t *testing.T) {
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ w.WriteHeader(http.StatusForbidden)
+ _, _ = w.Write([]byte(`AccessDenieddenied`))
+ }))
+ t.Cleanup(server.Close)
+
+ sess, err := session.NewSession(&aws.Config{
+ Region: aws.String("us-east-1"),
+ Endpoint: aws.String(server.URL),
+ Credentials: credentials.NewStaticCredentials("test", "test", ""),
+ MaxRetries: aws.Int(0),
+ })
+ require.NoError(t, err)
+
+ provider := &rdsProvider{
+ options: ProviderOptions{Id: "test"},
+ session: sess,
+ regions: &ec2.DescribeRegionsOutput{Regions: []*ec2.Region{{RegionName: aws.String("us-east-1")}, {RegionName: aws.String("eu-west-1")}}},
+ }
+ _, err = provider.GetResource(context.Background())
+ require.Error(t, err, "a denied listing in every region must not look like an empty account")
+}
+
+func TestListRDSResourcesKeepsTheCallThatSucceeded(t *testing.T) {
+ tests := []struct {
+ name string
+ deny string
+ wantDNS string
+ }{
+ {name: "clusters denied", deny: "DescribeDBClusters", wantDNS: "private-db.abc123.us-east-1.rds.amazonaws.com"},
+ {name: "instances denied", deny: "DescribeDBInstances", wantDNS: "aurora-cluster.cluster-abc123.us-east-1.rds.amazonaws.com"},
+ }
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ require.NoError(t, r.ParseForm())
+ if r.Form.Get("Action") == tt.deny {
+ w.WriteHeader(http.StatusForbidden)
+ _, _ = w.Write([]byte(`AccessDenieddenied`))
+ return
+ }
+ w.Header().Set("Content-Type", "text/xml")
+ switch r.Form.Get("Action") {
+ case "DescribeDBInstances":
+ _, _ = w.Write([]byte(describeDBInstancesPage2))
+ case "DescribeDBClusters":
+ _, _ = w.Write([]byte(describeDBClustersResponse))
+ default:
+ t.Errorf("unexpected action %q", r.Form.Get("Action"))
+ w.WriteHeader(http.StatusBadRequest)
+ }
+ }))
+ t.Cleanup(server.Close)
+
+ sess, err := session.NewSession(&aws.Config{
+ Region: aws.String("us-east-1"),
+ Endpoint: aws.String(server.URL),
+ Credentials: credentials.NewStaticCredentials("test", "test", ""),
+ MaxRetries: aws.Int(0),
+ })
+ require.NoError(t, err)
+
+ provider := &rdsProvider{
+ options: ProviderOptions{Id: "test"},
+ session: sess,
+ regions: &ec2.DescribeRegionsOutput{Regions: []*ec2.Region{{RegionName: aws.String("us-east-1")}}},
+ }
+ resources, err := provider.listRDSResources(context.Background(), rds.New(sess))
+ require.Error(t, err)
+ require.NotNil(t, resources)
+
+ var hosts []string
+ for _, item := range resources.Items {
+ hosts = append(hosts, item.DNSName)
+ }
+ require.Contains(t, hosts, tt.wantDNS)
+
+ kept, err := provider.GetResource(context.Background())
+ require.NoError(t, err)
+ var keptHosts []string
+ for _, item := range kept.Items {
+ keptHosts = append(keptHosts, item.DNSName)
+ }
+ require.Contains(t, keptHosts, tt.wantDNS)
+ })
+ }
+}