Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 17 additions & 1 deletion pkg/providers/aws/aws.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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
Expand Down Expand Up @@ -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
}
Expand Down Expand Up @@ -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"
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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
}
Expand Down
235 changes: 235 additions & 0 deletions pkg/providers/aws/rds.go
Original file line number Diff line number Diff line change
@@ -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
}
Loading
Loading