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
11 changes: 9 additions & 2 deletions cmd/datablue/handlers.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
4 changes: 2 additions & 2 deletions cmd/oceantv/broadcasthost/oceanmedia.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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
}
Expand Down
36 changes: 19 additions & 17 deletions cmd/oceantv/composite/store.go
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down
16 changes: 10 additions & 6 deletions storage/cloudflare.go
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@ import (
"context"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"fmt"
"time"

Expand Down Expand Up @@ -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{
Expand All @@ -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
}
Expand Down
144 changes: 57 additions & 87 deletions storage/cloudflare_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -28,34 +28,27 @@ 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.
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) {
Expand All @@ -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())
})
}
}

Expand Down Expand Up @@ -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")
}
})
}
Expand Down
Loading