diff --git a/tools/s3toss/main.go b/tools/s3toss/main.go index 37c800c7b5c3..1daeb94b9ba8 100755 --- a/tools/s3toss/main.go +++ b/tools/s3toss/main.go @@ -661,7 +661,7 @@ func migrateVnodes(mode string, vid, fid uint) { } } -func parseAccessString(as string) { +func parseAccessString(as string) error { items := strings.Split(as, ";") config.Endpoint = "" @@ -672,7 +672,13 @@ func parseAccessString(as string) { config.Region = "" for _, item := range items { + if item == "" { + continue + } part := strings.SplitN(item, "=", 2) + if len(part) != 2 { + return fmt.Errorf("invalid ssAccessString option %q: expected key=value", item) + } switch strings.ToLower(part[0]) { case "endpoint": config.Endpoint = part[1] @@ -688,6 +694,8 @@ func parseAccessString(as string) { config.Region = part[1] } } + + return nil } func parseTaosCfg(path string) error { @@ -724,24 +732,32 @@ func parseTaosCfg(path string) error { return err } } + if lvl < 0 || lvl >= len(config.DataDirs) { + return fmt.Errorf("invalid dataDir level %d: must be between 0 and %d", lvl, len(config.DataDirs)-1) + } config.DataDirs[lvl] = append(config.DataDirs[lvl], parts[1]) - case "ssAccessString": + case "ssaccessstring": if strings.HasPrefix(parts[1], "s3:") { ssConfig = true - parseAccessString(parts[1][3:]) + if err := parseAccessString(parts[1][3:]); err != nil { + return err + } } case "s3endpoint": if ssConfig || config.Endpoint != "" { continue } - if strings.HasPrefix(strings.ToLower(parts[1]), "http://") { + endpoint := strings.ToLower(parts[1]) + if strings.HasPrefix(endpoint, "http://") { config.Secure = false config.Endpoint = parts[1][7:] - } else { + } else if strings.HasPrefix(endpoint, "https://") { config.Secure = true config.Endpoint = parts[1][8:] + } else { + return fmt.Errorf("invalid s3endpoint %q: expected http:// or https:// prefix", parts[1]) } case "s3accesskey": diff --git a/tools/s3toss/main_test.go b/tools/s3toss/main_test.go new file mode 100644 index 000000000000..f0c3dbcd31f5 --- /dev/null +++ b/tools/s3toss/main_test.go @@ -0,0 +1,125 @@ +package main + +import ( + "os" + "path/filepath" + "strings" + "testing" +) + +func TestParseAccessStringAcceptsKnownAndUnknownOptions(t *testing.T) { + resetConfig() + + if err := parseAccessString("endpoint=s3.amazonaws.com;bucket=mybucket;uriStyle=path;protocol=http;accessKeyId=AK;secretAccessKey=SK;region=us-east-2"); err != nil { + t.Fatalf("parseAccessString returned error: %v", err) + } + + if config.Endpoint != "s3.amazonaws.com" { + t.Fatalf("Endpoint = %q, want s3.amazonaws.com", config.Endpoint) + } + if config.Bucket != "mybucket" { + t.Fatalf("Bucket = %q, want mybucket", config.Bucket) + } + if config.Secure { + t.Fatal("Secure = true, want false for protocol=http") + } + if config.AccessKey != "AK" || config.SecretKey != "SK" || config.Region != "us-east-2" { + t.Fatalf("unexpected credentials or region: access=%q secret=%q region=%q", config.AccessKey, config.SecretKey, config.Region) + } +} + +func TestParseAccessStringRejectsMalformedOption(t *testing.T) { + resetConfig() + + err := parseAccessString("endpoint") + if err == nil { + t.Fatal("parseAccessString succeeded for malformed option") + } + if !strings.Contains(err.Error(), "expected key=value") { + t.Fatalf("error = %q, want expected key=value", err) + } +} + +func TestParseTaosCfgParsesSsAccessString(t *testing.T) { + resetConfig() + + cfgPath := writeConfig(t, "dataDir /tmp/taos 0\nssAccessString s3:endpoint=s3.amazonaws.com;bucket=mybucket;uriStyle=path;protocol=http;accessKeyId=AK;secretAccessKey=SK;region=us-east-2\n") + + if err := parseTaosCfg(cfgPath); err != nil { + t.Fatalf("parseTaosCfg returned error: %v", err) + } + + if config.Endpoint != "s3.amazonaws.com" { + t.Fatalf("Endpoint = %q, want s3.amazonaws.com", config.Endpoint) + } + if config.Bucket != "mybucket" || config.AccessKey != "AK" || config.SecretKey != "SK" || config.Region != "us-east-2" { + t.Fatalf("unexpected S3 config: bucket=%q access=%q secret=%q region=%q", config.Bucket, config.AccessKey, config.SecretKey, config.Region) + } + if config.Secure { + t.Fatal("Secure = true, want false for protocol=http") + } +} + +func TestParseTaosCfgRejectsMalformedSsAccessString(t *testing.T) { + resetConfig() + + cfgPath := writeConfig(t, "dataDir /tmp/taos 0\nssAccessString s3:endpoint\n") + + err := parseTaosCfg(cfgPath) + if err == nil { + t.Fatal("parseTaosCfg succeeded for malformed ssAccessString") + } + if !strings.Contains(err.Error(), "expected key=value") { + t.Fatalf("error = %q, want expected key=value", err) + } +} + +func TestParseTaosCfgRejectsInvalidDataDirLevel(t *testing.T) { + resetConfig() + + cfgPath := writeConfig(t, "dataDir /tmp/taos 3\n") + + err := parseTaosCfg(cfgPath) + if err == nil { + t.Fatal("parseTaosCfg succeeded for invalid dataDir level") + } + if !strings.Contains(err.Error(), "invalid dataDir level 3") { + t.Fatalf("error = %q, want invalid dataDir level 3", err) + } +} + +func TestParseTaosCfgRejectsInvalidS3Endpoint(t *testing.T) { + resetConfig() + + cfgPath := writeConfig(t, "dataDir /tmp/taos 0\ns3endpoint localhost:9000\n") + + err := parseTaosCfg(cfgPath) + if err == nil { + t.Fatal("parseTaosCfg succeeded for invalid s3endpoint") + } + if !strings.Contains(err.Error(), "expected http:// or https:// prefix") { + t.Fatalf("error = %q, want endpoint prefix error", err) + } +} + +func writeConfig(t *testing.T, contents string) string { + t.Helper() + + cfgPath := filepath.Join(t.TempDir(), "taos.cfg") + if err := os.WriteFile(cfgPath, []byte(contents), 0o644); err != nil { + t.Fatalf("WriteFile failed: %v", err) + } + return cfgPath +} + +func resetConfig() { + config.BlockSize = 0 + config.DNode = 0 + config.DataDirs = [3][]string{} + config.Endpoint = "" + config.Secure = false + config.AccessKey = "" + config.SecretKey = "" + config.Bucket = "" + config.Region = "" +}