diff --git a/cmd/datablue/handlers.go b/cmd/datablue/handlers.go index cd704ae3..97298bb5 100644 --- a/cmd/datablue/handlers.go +++ b/cmd/datablue/handlers.go @@ -525,8 +525,15 @@ func varsHandler(w http.ResponseWriter, r *http.Request) { if v.IsSystemVariable() { continue } - resp += `"` + v.Name + `":"` + v.Value + `",` - + // Escape the value as a JSON string. Variable values may themselves + // contain JSON (e.g. storage configuration), which must be escaped in + // order to keep the enclosing response valid JSON. + value, err := json.Marshal(v.Value) + if err != nil { + writeError(w, fmt.Errorf("could not marshal variable %s: %w", v.Name, err)) + return + } + resp += `"` + v.Name + `":` + string(value) + `,` } vs := model.ComputeVarSum(vars) diff --git a/cmd/oceantv/broadcasthost/oceanmedia.go b/cmd/oceantv/broadcasthost/oceanmedia.go index bcaad6dd..5700eef3 100644 --- a/cmd/oceantv/broadcasthost/oceanmedia.go +++ b/cmd/oceantv/broadcasthost/oceanmedia.go @@ -63,7 +63,7 @@ func (o OceanMedia) New(args ...any) (any, error) { // If no arguments are provided, return an empty OceanMedia // so that we can still get the protocol. if len(args) == 0 { - return OceanMedia{}, nil + return &OceanMedia{}, nil } p, ok := args[0].(Params) if !ok { @@ -153,7 +153,7 @@ func (o *OceanMedia) BroadcastHealth(ctx context.Context, sid string) (string, e // This consists of a JSON encoded TempCredentials object. // streamName should be the broadcast event ID. func (o *OceanMedia) AuthKey(ctx context.Context, streamName string) (string, error) { - tempCreds, err := o.storageProvider.GenerateTempCredentials(ctx, 12*time.Hour, streamName) + tempCreds, err := o.storageProvider.GenerateTempCredentials(ctx, 12*time.Hour, streamName+"/") if err != nil { return "", err } diff --git a/cmd/oceantv/composite/store.go b/cmd/oceantv/composite/store.go index 14fcfd32..38de1d35 100644 --- a/cmd/oceantv/composite/store.go +++ b/cmd/oceantv/composite/store.go @@ -72,23 +72,25 @@ func AusOceanStore(settingsStore, mediaStore store) *Store { return NewStore( map[string]store{ - "Scalar": mediaStore, - "Text": mediaStore, - "MtsMedia": mediaStore, - "Device": settingsStore, - "Site": settingsStore, - "Signal": settingsStore, - "Notice": settingsStore, - "Trigger": settingsStore, - "Cron": settingsStore, - "Request": settingsStore, - "Variable": settingsStore, - "BinaryData": settingsStore, - "User": settingsStore, - "Sensor": settingsStore, - "SensorV2": settingsStore, - "Actuator": settingsStore, - "ActuatorV2": settingsStore, + "Scalar": mediaStore, + "Text": mediaStore, + "MtsMedia": mediaStore, + "BroadcastEvent": mediaStore, + "Notification": settingsStore, + "Device": settingsStore, + "Site": settingsStore, + "Signal": settingsStore, + "Notice": settingsStore, + "Trigger": settingsStore, + "Cron": settingsStore, + "Request": settingsStore, + "Variable": settingsStore, + "BinaryData": settingsStore, + "User": settingsStore, + "Sensor": settingsStore, + "SensorV2": settingsStore, + "Actuator": settingsStore, + "ActuatorV2": settingsStore, }, getKindFromQuery, ) diff --git a/storage/cloudflare.go b/storage/cloudflare.go index ecb7b4ce..1a5afbf5 100644 --- a/storage/cloudflare.go +++ b/storage/cloudflare.go @@ -28,6 +28,7 @@ import ( "context" "crypto/sha256" "encoding/base64" + "encoding/hex" "fmt" "time" @@ -56,11 +57,14 @@ func (c *Cloudflare) GenerateTempCredentials(ctx context.Context, ttl time.Durat now := time.Now() claims := map[string]interface{}{ - "exp": now.Add(ttl).Unix(), - "iat": now.Unix(), - "sub": c.accountID, - "aud": fmt.Sprintf("%s.r2.cloudflarestorage.com", c.accountID), - "bucket": c.bucket, + "exp": now.Add(ttl).Unix(), + "iat": now.Unix(), + "iss": c.accessKey, + "sub": c.accountID, + "aud": fmt.Sprintf("%s.r2.cloudflarestorage.com", c.accountID), + "bucket": c.bucket, + "scope": "object-read-write", + "ttlSeconds": int64(ttl.Seconds()), } if prefix != "" { claims["paths"] = map[string][]string{ @@ -80,7 +84,7 @@ func (c *Cloudflare) GenerateTempCredentials(ctx context.Context, ttl time.Durat return &TempCredentials{ AccessKey: c.accessKey, - SecretKey: string(secretKey[:]), + SecretKey: hex.EncodeToString(secretKey[:]), SessionToken: sessionToken, }, nil } diff --git a/storage/cloudflare_test.go b/storage/cloudflare_test.go index 5a811f77..897060ad 100644 --- a/storage/cloudflare_test.go +++ b/storage/cloudflare_test.go @@ -28,12 +28,15 @@ import ( "context" "crypto/sha256" "encoding/base64" + "encoding/hex" "fmt" "strings" "testing" "time" "github.com/ausocean/cloud/gauth" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" ) // Ensure Cloudflare implements Provider interface. @@ -41,21 +44,11 @@ var _ Provider = (*Cloudflare)(nil) func TestNewCloudflare(t *testing.T) { cf := NewCloudflare("test-account", "test-access", "test-secret", "test-bucket") - if cf == nil { - t.Fatal("NewCloudflare returned nil") - } - if cf.accountID != "test-account" { - t.Errorf("got accountID %q, want %q", cf.accountID, "test-account") - } - if cf.accessKey != "test-access" { - t.Errorf("got accessKey %q, want %q", cf.accessKey, "test-access") - } - if cf.secretKey != "test-secret" { - t.Errorf("got secretKey %q, want %q", cf.secretKey, "test-secret") - } - if cf.bucket != "test-bucket" { - t.Errorf("got bucket %q, want %q", cf.bucket, "test-bucket") - } + require.NotNil(t, cf) + assert.Equal(t, "test-account", cf.accountID) + assert.Equal(t, "test-access", cf.accessKey) + assert.Equal(t, "test-secret", cf.secretKey) + assert.Equal(t, "test-bucket", cf.bucket) } func TestGetBaseURL(t *testing.T) { @@ -74,11 +67,10 @@ func TestGetBaseURL(t *testing.T) { } for _, tt := range tests { - cf := NewCloudflare(tt.accountID, "access", "secret", "bucket") - got := cf.GetBaseURL() - if got != tt.want { - t.Errorf("GetBaseURL() = %q, want %q", got, tt.want) - } + t.Run(tt.accountID, func(t *testing.T) { + cf := NewCloudflare(tt.accountID, "access", "secret", "bucket") + assert.Equal(t, tt.want, cf.GetBaseURL()) + }) } } @@ -132,93 +124,71 @@ func TestGenerateTempCredentials(t *testing.T) { before := time.Now().Unix() creds, err := cf.GenerateTempCredentials(context.Background(), tt.ttl, tt.prefix) after := time.Now().Unix() - if err != nil { - t.Fatalf("GenerateTempCredentials() error = %v", err) - } - if creds == nil { - t.Fatal("GenerateTempCredentials() returned nil credentials") - } - // Verify AccessKey - if creds.AccessKey != accessKey { - t.Errorf("got AccessKey %q, want %q", creds.AccessKey, accessKey) - } + require.NoError(t, err) + require.NotNil(t, creds) - // Verify SessionToken is base64 encoded and begins with "jwt/" + // Verify AccessKey. + assert.Equal(t, accessKey, creds.AccessKey) + + // Verify SessionToken is base64 encoded and begins with "jwt/". decodedBytes, err := base64.StdEncoding.DecodeString(creds.SessionToken) - if err != nil { - t.Fatalf("SessionToken is not valid base64: %v", err) - } + require.NoError(t, err, "SessionToken is not valid base64") decodedToken := string(decodedBytes) - if !strings.HasPrefix(decodedToken, "jwt/") { - t.Fatalf("decoded SessionToken %q does not have prefix 'jwt/'", decodedToken) - } + require.True(t, strings.HasPrefix(decodedToken, "jwt/"), "decoded SessionToken %q does not have prefix 'jwt/'", decodedToken) jwtStr := strings.TrimPrefix(decodedToken, "jwt/") - // Verify SecretKey is the sha256 sum of the JWT string + // Verify SecretKey is the hex-encoded sha256 sum of the JWT string. expectedSecretKeyBytes := sha256.Sum256([]byte(jwtStr)) - if creds.SecretKey != string(expectedSecretKeyBytes[:]) { - t.Errorf("SecretKey does not match sha256 of JWT") - } + assert.Equal(t, hex.EncodeToString(expectedSecretKeyBytes[:]), creds.SecretKey) - // Verify JWT claims by parsing with gauth.GetClaims using the secret key + // Verify JWT claims by parsing with gauth.GetClaims using the secret key. claims, err := gauth.GetClaims(jwtStr, []byte(secretKey)) - if err != nil { - t.Fatalf("gauth.GetClaims() failed: %v", err) - } + require.NoError(t, err) - // Check sub - if sub, ok := claims["sub"].(string); !ok || sub != accountID { - t.Errorf("claims[sub] = %v, want %q", claims["sub"], accountID) - } + // Check iss. + assert.Equal(t, accessKey, claims["iss"]) - // Check aud - wantAud := fmt.Sprintf("%s.r2.cloudflarestorage.com", accountID) - if aud, ok := claims["aud"].(string); !ok || aud != wantAud { - t.Errorf("claims[aud] = %v, want %q", claims["aud"], wantAud) - } + // Check sub. + assert.Equal(t, accountID, claims["sub"]) - // Check bucket - if b, ok := claims["bucket"].(string); !ok || b != bucket { - t.Errorf("claims[bucket] = %v, want %q", claims["bucket"], bucket) - } + // Check aud. + assert.Equal(t, fmt.Sprintf("%s.r2.cloudflarestorage.com", accountID), claims["aud"]) - // Check iat and exp - iatVal, ok := claims["iat"].(float64) - if !ok { - t.Fatalf("claims[iat] is not a number: %v", claims["iat"]) - } - iat := int64(iatVal) - if iat < before || iat > after { - t.Errorf("claims[iat] %d out of expected range [%d, %d]", iat, before, after) - } + // Check bucket. + assert.Equal(t, bucket, claims["bucket"]) - expVal, ok := claims["exp"].(float64) - if !ok { - t.Fatalf("claims[exp] is not a number: %v", claims["exp"]) - } - exp := int64(expVal) - expectedExpMin := before + int64(tt.wantTTL.Seconds()) - expectedExpMax := after + int64(tt.wantTTL.Seconds()) - if exp < expectedExpMin || exp > expectedExpMax { - t.Errorf("claims[exp] %d out of expected range [%d, %d]", exp, expectedExpMin, expectedExpMax) - } + // Check scope. + assert.Equal(t, "object-read-write", claims["scope"]) + + // Check iat. + iat, ok := claims["iat"].(float64) + require.True(t, ok, "claims[iat] is not a number: %v", claims["iat"]) + assert.GreaterOrEqual(t, int64(iat), before) + assert.LessOrEqual(t, int64(iat), after) + + // Check exp. + exp, ok := claims["exp"].(float64) + require.True(t, ok, "claims[exp] is not a number: %v", claims["exp"]) + assert.GreaterOrEqual(t, int64(exp), before+int64(tt.wantTTL.Seconds())) + assert.LessOrEqual(t, int64(exp), after+int64(tt.wantTTL.Seconds())) + + // Check ttlSeconds. + ttlSeconds, ok := claims["ttlSeconds"].(float64) + require.True(t, ok, "claims[ttlSeconds] is not a number: %v", claims["ttlSeconds"]) + assert.Equal(t, tt.wantTTL.Seconds(), ttlSeconds) - // Check prefix / paths + // Check prefix / paths. if tt.wantPrefix != "" { paths, ok := claims["paths"].(map[string]interface{}) - if !ok { - t.Fatalf("claims[paths] missing or invalid type: %v", claims["paths"]) - } + require.True(t, ok, "claims[paths] missing or invalid type: %v", claims["paths"]) prefixPaths, ok := paths["prefixPaths"].([]interface{}) - if !ok || len(prefixPaths) != 1 || prefixPaths[0] != tt.wantPrefix { - t.Errorf("claims[paths][prefixPaths] = %v, want [%q]", paths["prefixPaths"], tt.wantPrefix) - } + require.True(t, ok, "claims[paths][prefixPaths] missing or invalid type: %v", paths["prefixPaths"]) + require.Len(t, prefixPaths, 1) + assert.Equal(t, tt.wantPrefix, prefixPaths[0]) } else { - if _, exists := claims["paths"]; exists { - t.Errorf("claims[paths] should not exist when prefix is empty, got %v", claims["paths"]) - } + assert.NotContains(t, claims, "paths") } }) }