Skip to content

Commit 727c2b7

Browse files
committed
feat(volcengine): add DNS support
1 parent 352336c commit 727c2b7

7 files changed

Lines changed: 310 additions & 9 deletions

File tree

README.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -58,7 +58,7 @@ CloudToolKit is intended to help defenders verify:
5858
| AWS | EC2, S3, IAM | iam-user-check, bucket-check |
5959
| Azure | Virtual Machines, Blob Storage | - |
6060
| GCP | Compute Engine, Cloud DNS, IAM | - |
61-
| Volcengine | ECS, IAM, TOS, RDS | iam-user-check, bucket-check, instance-cmd-check |
61+
| Volcengine | ECS, IAM, TOS, RDS, DNS | iam-user-check, bucket-check, instance-cmd-check |
6262
| JDCloud | VM, IAM, OSS | - |
6363

6464
## Quick Start

pkg/providers/volcengine/api/endpoint.go

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@ const (
99

1010
var globalServices = map[string]struct{}{
1111
"billing": {},
12+
"dns": {},
1213
"iam": {},
1314
}
1415

pkg/providers/volcengine/api/endpoint_test.go

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@ func TestResolveEndpoint(t *testing.T) {
1212
}{
1313
{name: "global iam", service: "iam", region: "cn-beijing", want: "https://iam.volcengineapi.com"},
1414
{name: "global billing", service: "billing", region: "", want: "https://billing.volcengineapi.com"},
15+
{name: "global dns", service: "dns", region: "cn-beijing", want: "https://dns.volcengineapi.com"},
1516
{name: "regional ecs", service: "ecs", region: "cn-shanghai", want: "https://ecs.cn-shanghai.volcengineapi.com"},
1617
{name: "regional rds mysql alias", service: "rds_mysql", region: "cn-beijing", want: "https://rds-mysql.cn-beijing.volcengineapi.com"},
1718
{name: "regional rds postgresql alias", service: "rds_postgresql", region: "cn-guangzhou", want: "https://rds-postgresql.cn-guangzhou.volcengineapi.com"},
Lines changed: 80 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,80 @@
1+
package api
2+
3+
import (
4+
"context"
5+
"encoding/json"
6+
"net/http"
7+
)
8+
9+
const dnsAPIVersion = "2018-08-01"
10+
11+
type ListDNSZonesResponse struct {
12+
Total int32 `json:"Total"`
13+
Zones []DNSZone `json:"Zones"`
14+
}
15+
16+
type DNSZone struct {
17+
ZID int64 `json:"ZID"`
18+
ZoneName string `json:"ZoneName"`
19+
}
20+
21+
type ListDNSRecordsResponse struct {
22+
PageNumber int32 `json:"PageNumber"`
23+
PageSize int32 `json:"PageSize"`
24+
Records []DNSRecord `json:"Records"`
25+
TotalCount int32 `json:"TotalCount"`
26+
}
27+
28+
type DNSRecord struct {
29+
Enable *bool `json:"Enable,omitempty"`
30+
FQDN string `json:"FQDN"`
31+
Host string `json:"Host"`
32+
Type string `json:"Type"`
33+
Value string `json:"Value"`
34+
}
35+
36+
type listDNSZonesInput struct {
37+
PageNumber int32 `json:"PageNumber"`
38+
PageSize int32 `json:"PageSize"`
39+
}
40+
41+
type listDNSRecordsInput struct {
42+
PageNumber int32 `json:"PageNumber"`
43+
PageSize int32 `json:"PageSize"`
44+
ZID int64 `json:"ZID"`
45+
}
46+
47+
func (c *Client) ListDNSZones(ctx context.Context, pageNumber, pageSize int32) (ListDNSZonesResponse, error) {
48+
var out ListDNSZonesResponse
49+
err := c.doDNSAction(ctx, "ListZones", listDNSZonesInput{
50+
PageNumber: pageNumber,
51+
PageSize: pageSize,
52+
}, &out)
53+
return out, err
54+
}
55+
56+
func (c *Client) ListDNSRecords(ctx context.Context, zid int64, pageNumber, pageSize int32) (ListDNSRecordsResponse, error) {
57+
var out ListDNSRecordsResponse
58+
err := c.doDNSAction(ctx, "ListRecords", listDNSRecordsInput{
59+
PageNumber: pageNumber,
60+
PageSize: pageSize,
61+
ZID: zid,
62+
}, &out)
63+
return out, err
64+
}
65+
66+
func (c *Client) doDNSAction(ctx context.Context, action string, payload any, out any) error {
67+
body, err := json.Marshal(payload)
68+
if err != nil {
69+
return err
70+
}
71+
return c.DoOpenAPI(ctx, Request{
72+
Service: "dns",
73+
Version: dnsAPIVersion,
74+
Action: action,
75+
Method: http.MethodPost,
76+
Path: "/",
77+
Body: body,
78+
Idempotent: true,
79+
}, out)
80+
}
Lines changed: 130 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,130 @@
1+
package dns
2+
3+
import (
4+
"context"
5+
"errors"
6+
"strings"
7+
8+
"github.com/404tk/cloudtoolkit/pkg/providers/volcengine/api"
9+
"github.com/404tk/cloudtoolkit/pkg/runtime/paginate"
10+
"github.com/404tk/cloudtoolkit/pkg/schema"
11+
"github.com/404tk/cloudtoolkit/utils/logger"
12+
)
13+
14+
const pageSize int32 = 100
15+
16+
type Driver struct {
17+
Client *api.Client
18+
}
19+
20+
var errNilAPIClient = errors.New("volcengine dns: nil api client")
21+
22+
func (d *Driver) GetDomains(ctx context.Context) ([]schema.Domain, error) {
23+
list := []schema.Domain{}
24+
select {
25+
case <-ctx.Done():
26+
return list, nil
27+
default:
28+
logger.Info("List DNS ...")
29+
}
30+
if d.Client == nil {
31+
return list, errNilAPIClient
32+
}
33+
34+
zones, err := paginate.Fetch[api.DNSZone, int32](ctx, func(ctx context.Context, pageNumber int32) (paginate.Page[api.DNSZone, int32], error) {
35+
if pageNumber == 0 {
36+
pageNumber = 1
37+
}
38+
resp, err := d.Client.ListDNSZones(ctx, pageNumber, pageSize)
39+
if err != nil {
40+
logger.Error("List zones failed.")
41+
return paginate.Page[api.DNSZone, int32]{}, err
42+
}
43+
return paginate.Page[api.DNSZone, int32]{
44+
Items: resp.Zones,
45+
Next: pageNumber + 1,
46+
Done: pageDone(pageNumber, pageSize, resp.Total, len(resp.Zones)),
47+
}, nil
48+
})
49+
if err != nil {
50+
return list, err
51+
}
52+
53+
for _, zone := range zones {
54+
name := strings.TrimSpace(zone.ZoneName)
55+
if name == "" || zone.ZID == 0 {
56+
continue
57+
}
58+
records, err := d.listRecords(ctx, zone.ZID)
59+
if err != nil {
60+
logger.Error("List records failed.")
61+
return list, err
62+
}
63+
list = append(list, schema.Domain{
64+
DomainName: name,
65+
Records: records,
66+
})
67+
}
68+
69+
return list, nil
70+
}
71+
72+
func (d *Driver) listRecords(ctx context.Context, zid int64) ([]schema.Record, error) {
73+
records, err := paginate.Fetch[api.DNSRecord, int32](ctx, func(ctx context.Context, pageNumber int32) (paginate.Page[api.DNSRecord, int32], error) {
74+
if pageNumber == 0 {
75+
pageNumber = 1
76+
}
77+
resp, err := d.Client.ListDNSRecords(ctx, zid, pageNumber, pageSize)
78+
if err != nil {
79+
return paginate.Page[api.DNSRecord, int32]{}, err
80+
}
81+
return paginate.Page[api.DNSRecord, int32]{
82+
Items: resp.Records,
83+
Next: pageNumber + 1,
84+
Done: pageDone(pageNumber, pageSize, resp.TotalCount, len(resp.Records)),
85+
}, nil
86+
})
87+
if err != nil {
88+
return nil, err
89+
}
90+
91+
list := make([]schema.Record, 0, len(records))
92+
for _, record := range records {
93+
list = append(list, schema.Record{
94+
RR: firstNonEmpty(strings.TrimSpace(record.Host), strings.TrimSpace(record.FQDN), "@"),
95+
Type: strings.TrimSpace(record.Type),
96+
Value: strings.TrimSpace(record.Value),
97+
Status: recordStatus(record.Enable),
98+
})
99+
}
100+
return list, nil
101+
}
102+
103+
func recordStatus(enable *bool) string {
104+
if enable == nil {
105+
return ""
106+
}
107+
if *enable {
108+
return "ENABLE"
109+
}
110+
return "DISABLE"
111+
}
112+
113+
func pageDone(pageNumber, pageSize, total int32, items int) bool {
114+
if items == 0 {
115+
return true
116+
}
117+
if total <= 0 {
118+
return int32(items) < pageSize
119+
}
120+
return pageNumber*pageSize >= total
121+
}
122+
123+
func firstNonEmpty(values ...string) string {
124+
for _, value := range values {
125+
if trimmed := strings.TrimSpace(value); trimmed != "" {
126+
return trimmed
127+
}
128+
}
129+
return ""
130+
}

pkg/providers/volcengine/volcengine.go

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@ import (
99
_api "github.com/404tk/cloudtoolkit/pkg/providers/volcengine/api"
1010
_auth "github.com/404tk/cloudtoolkit/pkg/providers/volcengine/auth"
1111
"github.com/404tk/cloudtoolkit/pkg/providers/volcengine/billing"
12+
_dns "github.com/404tk/cloudtoolkit/pkg/providers/volcengine/dns"
1213
"github.com/404tk/cloudtoolkit/pkg/providers/volcengine/ecs"
1314
"github.com/404tk/cloudtoolkit/pkg/providers/volcengine/iam"
1415
"github.com/404tk/cloudtoolkit/pkg/providers/volcengine/rds"
@@ -74,6 +75,10 @@ func (p *Provider) Resources(ctx context.Context) (schema.Resources, error) {
7475
schema.AppendAssets(&list, hosts)
7576
list.AddError("host", err)
7677
case "domain":
78+
d := &_dns.Driver{Client: p.apiClient}
79+
domains, err := d.GetDomains(ctx)
80+
schema.AppendAssets(&list, domains)
81+
list.AddError("domain", err)
7782
case "account":
7883
d := &iam.Driver{Client: p.apiClient, Region: p.region}
7984
users, err := d.ListUsers(ctx)

pkg/providers/volcengine/volcengine_test.go

Lines changed: 92 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -125,6 +125,84 @@ func TestProviderResourcesDatabaseUsesRDSDrivers(t *testing.T) {
125125
}
126126
}
127127

128+
func TestProviderResourcesDomainUsesDNSDriver(t *testing.T) {
129+
restoreCloudlist := setCloudlist([]string{"domain"})
130+
defer restoreCloudlist()
131+
132+
logger.SetOutput(io.Discard)
133+
t.Cleanup(func() {
134+
logger.SetOutput(nil)
135+
})
136+
137+
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
138+
values, err := url.ParseQuery(r.URL.RawQuery)
139+
if err != nil {
140+
t.Fatalf("ParseQuery() error = %v", err)
141+
}
142+
if service := signedService(t, r); service != "dns" {
143+
t.Fatalf("unexpected service: %s", service)
144+
}
145+
switch values.Get("Action") {
146+
case "ListZones":
147+
assertJSONBody(t, r, map[string]any{
148+
"PageNumber": float64(1),
149+
"PageSize": float64(100),
150+
})
151+
_, _ = w.Write([]byte(`{"Total":2,"Zones":[{"ZoneName":"example.com","ZID":101},{"ZoneName":"example.org","ZID":102}]}`))
152+
case "ListRecords":
153+
body := decodeJSONBody(t, r)
154+
if body["PageNumber"] != float64(1) || body["PageSize"] != float64(100) {
155+
t.Fatalf("unexpected ListRecords body: %v", body)
156+
}
157+
switch body["ZID"] {
158+
case float64(101):
159+
_, _ = w.Write([]byte(`{"PageNumber":1,"PageSize":100,"TotalCount":2,"Records":[{"Host":"@","Type":"A","Value":"1.1.1.1","Enable":true},{"Host":"www","Type":"CNAME","Value":"target.example.com","Enable":false}]}`))
160+
case float64(102):
161+
_, _ = w.Write([]byte(`{"PageNumber":1,"PageSize":100,"TotalCount":1,"Records":[{"Host":"api","Type":"A","Value":"2.2.2.2","Enable":true}]}`))
162+
default:
163+
t.Fatalf("unexpected ZID: %v", body["ZID"])
164+
}
165+
default:
166+
t.Fatalf("unexpected action: %s", values.Get("Action"))
167+
}
168+
}))
169+
defer server.Close()
170+
171+
provider, err := newProvider(testOptions(map[string]string{
172+
utils.Provider: "volcengine",
173+
utils.Region: "cn-guangzhou",
174+
}), testClientOptions(server.URL)...)
175+
if err != nil {
176+
t.Fatalf("newProvider() error = %v", err)
177+
}
178+
179+
resources, err := provider.Resources(context.Background())
180+
if err != nil {
181+
t.Fatalf("Resources() error = %v", err)
182+
}
183+
if len(resources.Errors) != 0 {
184+
t.Fatalf("unexpected resource errors: %+v", resources.Errors)
185+
}
186+
187+
got := map[string]schema.Domain{}
188+
for _, asset := range resources.Assets {
189+
domain, ok := asset.(schema.Domain)
190+
if !ok {
191+
t.Fatalf("unexpected asset type: %T", asset)
192+
}
193+
got[domain.DomainName] = domain
194+
}
195+
if len(got) != 2 {
196+
t.Fatalf("unexpected domain count: %d", len(got))
197+
}
198+
if len(got["example.com"].Records) != 2 || got["example.com"].Records[1].Status != "DISABLE" {
199+
t.Fatalf("unexpected example.com records: %+v", got["example.com"])
200+
}
201+
if len(got["example.org"].Records) != 1 || got["example.org"].Records[0].Value != "2.2.2.2" {
202+
t.Fatalf("unexpected example.org records: %+v", got["example.org"])
203+
}
204+
}
205+
128206
func testOptions(overrides map[string]string) schema.Options {
129207
options := schema.Options{
130208
utils.AccessKey: "AKID",
@@ -156,6 +234,19 @@ func setCloudlist(values []string) func() {
156234
}
157235

158236
func assertJSONBody(t *testing.T, r *http.Request, want map[string]any) {
237+
t.Helper()
238+
got := decodeJSONBody(t, r)
239+
if len(got) != len(want) {
240+
t.Fatalf("unexpected body map length: got=%v want=%v", got, want)
241+
}
242+
for key, wantValue := range want {
243+
if got[key] != wantValue {
244+
t.Fatalf("unexpected body field %s: got=%v want=%v", key, got[key], wantValue)
245+
}
246+
}
247+
}
248+
249+
func decodeJSONBody(t *testing.T, r *http.Request) map[string]any {
159250
t.Helper()
160251
defer r.Body.Close()
161252
body, err := io.ReadAll(r.Body)
@@ -166,14 +257,7 @@ func assertJSONBody(t *testing.T, r *http.Request, want map[string]any) {
166257
if err := json.Unmarshal(body, &got); err != nil {
167258
t.Fatalf("Unmarshal() error = %v body=%s", err, string(body))
168259
}
169-
if len(got) != len(want) {
170-
t.Fatalf("unexpected body map length: got=%v want=%v", got, want)
171-
}
172-
for key, wantValue := range want {
173-
if got[key] != wantValue {
174-
t.Fatalf("unexpected body field %s: got=%v want=%v", key, got[key], wantValue)
175-
}
176-
}
260+
return got
177261
}
178262

179263
func signedService(t *testing.T, r *http.Request) string {

0 commit comments

Comments
 (0)