diff --git a/cmd/appview/serve.go b/cmd/appview/serve.go index eeeb00b..5c0bd12 100644 --- a/cmd/appview/serve.go +++ b/cmd/appview/serve.go @@ -9,14 +9,12 @@ import ( "net/http" "os" "os/signal" - "path/filepath" "syscall" "time" "github.com/distribution/distribution/v3/configuration" "github.com/distribution/distribution/v3/registry" "github.com/distribution/distribution/v3/registry/handlers" - sqlite3 "github.com/mattn/go-sqlite3" "github.com/spf13/cobra" "atcr.io/pkg/appview/middleware" @@ -32,34 +30,6 @@ import ( "github.com/gorilla/mux" ) -// Define sensitive tables that should never be accessible from public queries -var sensitiveTables = map[string]bool{ - "oauth_sessions": true, // OAuth tokens - "ui_sessions": true, // Session IDs - "oauth_auth_requests": true, // OAuth state - "devices": true, // Device secret hashes - "pending_device_auth": true, // Pending device secrets -} - -// readOnlyAuthorizerCallback blocks access to sensitive tables -func readOnlyAuthorizerCallback(action int, arg1, arg2, dbName string) int { - // arg1 contains the table name for most operations - tableName := arg1 - - // Block any access to sensitive tables - if action == sqlite3.SQLITE_READ || action == sqlite3.SQLITE_UPDATE || - action == sqlite3.SQLITE_INSERT || action == sqlite3.SQLITE_DELETE || - action == sqlite3.SQLITE_SELECT { - if sensitiveTables[tableName] { - fmt.Printf("SECURITY: Blocked access to sensitive table '%s' (action=%d)\n", tableName, action) - return sqlite3.SQLITE_DENY - } - } - - // Allow everything else - return sqlite3.SQLITE_OK -} - var serveCmd = &cobra.Command{ Use: "serve", Short: "Start the ATCR registry server", @@ -72,15 +42,6 @@ See .env.appview.example for available environment variables.`, } func init() { - // Register a custom SQLite driver with authorizer for read-only public queries - sql.Register("sqlite3_readonly_public", - &sqlite3.SQLiteDriver{ - ConnectHook: func(conn *sqlite3.SQLiteConn) error { - conn.RegisterAuthorizer(readOnlyAuthorizerCallback) - return nil - }, - }) - // Replace the default serve command with our custom one for i, cmd := range registry.RootCmd.Commands() { if cmd.Name() == "serve" { @@ -93,7 +54,7 @@ func init() { func serveRegistry(cmd *cobra.Command, args []string) error { // Load configuration from environment variables fmt.Println("Loading configuration from environment variables...") - config, err := loadConfigFromEnv() + config, err := appview.LoadConfigFromEnv() if err != nil { return fmt.Errorf("failed to load config from environment: %w", err) } @@ -101,7 +62,12 @@ func serveRegistry(cmd *cobra.Command, args []string) error { // Initialize UI database first (required for all stores) fmt.Println("Initializing UI database...") - uiDatabase, uiReadOnlyDB, uiSessionStore := initializeDatabase() + uiEnabled := os.Getenv("ATCR_UI_ENABLED") != "false" + dbPath := os.Getenv("ATCR_UI_DATABASE_PATH") + if dbPath == "" { + dbPath = "/var/lib/atcr/ui.db" + } + uiDatabase, uiReadOnlyDB, uiSessionStore := db.InitializeDatabase(uiEnabled, dbPath) if uiDatabase == nil { return fmt.Errorf("failed to initialize UI database - required for session storage") } @@ -158,7 +124,7 @@ func serveRegistry(cmd *cobra.Command, args []string) error { // Expected format: "did:web:hold01.atcr.io" // To find a hold's DID, visit: https://hold01.atcr.io/.well-known/did.json // The extraction function normalizes URLs to DIDs for consistency - defaultHoldDID := extractDefaultHoldDID(config) + defaultHoldDID := appview.ExtractDefaultHoldDID(config) // Initialize UI routes with OAuth app, refresher, and device store uiTemplates, uiRouter := initializeUIRoutes(uiDatabase, uiReadOnlyDB, uiSessionStore, oauthApp, refresher, baseURL, deviceStore, defaultHoldDID) @@ -304,10 +270,10 @@ func initializeAuthKeys(config *configuration.Configuration) error { return nil } - privateKeyPath := getStringParam(tokenParams, "privatekey", "/var/lib/atcr/auth/private-key.pem") - issuerName := getStringParam(tokenParams, "issuer", "atcr.io") - service := getStringParam(tokenParams, "service", "atcr.io") - expirationSecs := getIntParam(tokenParams, "expiration", 300) + privateKeyPath := appview.GetStringParam(tokenParams, "privatekey", "/var/lib/atcr/auth/private-key.pem") + issuerName := appview.GetStringParam(tokenParams, "issuer", "atcr.io") + service := appview.GetStringParam(tokenParams, "service", "atcr.io") + expirationSecs := appview.GetIntParam(tokenParams, "expiration", 300) // Create issuer (this will generate the key if it doesn't exist) _, err := token.NewIssuer( @@ -331,10 +297,10 @@ func createTokenIssuer(config *configuration.Configuration) (*token.Issuer, erro return nil, fmt.Errorf("token auth not configured") } - privateKeyPath := getStringParam(tokenParams, "privatekey", "/var/lib/atcr/auth/private-key.pem") - issuerName := getStringParam(tokenParams, "issuer", "atcr.io") - service := getStringParam(tokenParams, "service", "atcr.io") - expirationSecs := getIntParam(tokenParams, "expiration", 300) + privateKeyPath := appview.GetStringParam(tokenParams, "privatekey", "/var/lib/atcr/auth/private-key.pem") + issuerName := appview.GetStringParam(tokenParams, "issuer", "atcr.io") + service := appview.GetStringParam(tokenParams, "service", "atcr.io") + expirationSecs := appview.GetIntParam(tokenParams, "expiration", 300) return token.NewIssuer( privateKeyPath, @@ -344,120 +310,6 @@ func createTokenIssuer(config *configuration.Configuration) (*token.Issuer, erro ) } -// Helper functions to extract values from config parameters -func getStringParam(params configuration.Parameters, key, defaultValue string) string { - if v, ok := params[key]; ok { - if s, ok := v.(string); ok { - return s - } - } - return defaultValue -} - -func getIntParam(params configuration.Parameters, key string, defaultValue int) int { - if v, ok := params[key]; ok { - if i, ok := v.(int); ok { - return i - } - } - return defaultValue -} - -// extractDefaultHoldDID extracts the default hold DID from middleware config -// Returns a DID (e.g., "did:web:hold01.atcr.io") -// To find a hold's DID, visit: https://hold-url/.well-known/did.json -func extractDefaultHoldDID(config *configuration.Configuration) string { - // Navigate through: middleware.registry[].options.default_hold_did - registryMiddleware, ok := config.Middleware["registry"] - if !ok { - return "" - } - - // Find atproto-resolver middleware - for _, mw := range registryMiddleware { - // Check if this is the atproto-resolver - if mw.Name != "atproto-resolver" { - continue - } - - // Extract options - options is configuration.Parameters which is map[string]any - if mw.Options != nil { - if holdDID, ok := mw.Options["default_hold_did"].(string); ok { - return holdDID - } - } - } - - return "" -} - -// initializeDatabase initializes the SQLite database and session store -// Returns: (read-write DB, read-only DB, session store) -func initializeDatabase() (*sql.DB, *sql.DB, *db.SessionStore) { - // Check if UI is enabled (optional configuration) - uiEnabled := os.Getenv("ATCR_UI_ENABLED") - if uiEnabled == "false" { - return nil, nil, nil - } - - // Get database path - dbPath := os.Getenv("ATCR_UI_DATABASE_PATH") - if dbPath == "" { - dbPath = "/var/lib/atcr/ui.db" - } - - // Ensure directory exists - dbDir := filepath.Dir(dbPath) - if err := os.MkdirAll(dbDir, 0700); err != nil { - fmt.Printf("Warning: Failed to create UI database directory: %v\n", err) - return nil, nil, nil - } - - // Initialize read-write database (for writes and auth operations) - database, err := db.InitDB(dbPath) - if err != nil { - fmt.Printf("Warning: Failed to initialize UI database: %v\n", err) - return nil, nil, nil - } - - // Open read-only connection for public queries (search, user pages, etc.) - // Uses custom driver with SQLite authorizer that blocks sensitive tables - // This prevents accidental writes and blocks access to sensitive tables even if SQL injection occurs - readOnlyDB, err := sql.Open("sqlite3_readonly_public", "file:"+dbPath+"?mode=ro") - if err != nil { - fmt.Printf("Warning: Failed to open read-only database connection: %v\n", err) - return nil, nil, nil - } - - fmt.Printf("UI database (readonly) initialized at %s\n", dbPath) - - // Create SQLite-backed session store - sessionStore := db.NewSessionStore(database) - - // Start cleanup goroutines for all SQLite stores - go func() { - ticker := time.NewTicker(5 * time.Minute) - defer ticker.Stop() - for range ticker.C { - ctx := context.Background() - - // Cleanup UI sessions - sessionStore.Cleanup() - - // Cleanup OAuth sessions (older than 30 days) - oauthStore := db.NewOAuthStore(database) - oauthStore.CleanupOldSessions(ctx, 30*24*time.Hour) - oauthStore.CleanupExpiredAuthRequests(ctx) - - // Cleanup device pending auths - deviceStore := db.NewDeviceStore(database) - deviceStore.CleanupExpired() - } - }() - - return database, readOnlyDB, sessionStore -} - // initializeUIRoutes initializes the web UI routes // database: read-write connection for auth and writes // readOnlyDB: read-only connection for public queries (search, user pages, etc.) diff --git a/cmd/appview/config.go b/pkg/appview/config.go similarity index 72% rename from cmd/appview/config.go rename to pkg/appview/config.go index 9097f8b..c920414 100644 --- a/cmd/appview/config.go +++ b/pkg/appview/config.go @@ -1,4 +1,4 @@ -package main +package appview import ( "crypto/rand" @@ -12,9 +12,9 @@ import ( "github.com/distribution/distribution/v3/configuration" ) -// loadConfigFromEnv builds a complete configuration from environment variables +// LoadConfigFromEnv builds a complete configuration from environment variables // This follows the same pattern as the hold service (no config files, only env vars) -func loadConfigFromEnv() (*configuration.Configuration, error) { +func LoadConfigFromEnv() (*configuration.Configuration, error) { config := &configuration.Configuration{} // Version @@ -56,8 +56,8 @@ func loadConfigFromEnv() (*configuration.Configuration, error) { // buildLogConfig creates logging configuration from environment variables func buildLogConfig() configuration.Log { - level := getEnvOrDefault("ATCR_LOG_LEVEL", "info") - formatter := getEnvOrDefault("ATCR_LOG_FORMATTER", "text") + level := GetEnvOrDefault("ATCR_LOG_LEVEL", "info") + formatter := GetEnvOrDefault("ATCR_LOG_FORMATTER", "text") return configuration.Log{ Level: configuration.Loglevel(level), @@ -70,8 +70,8 @@ func buildLogConfig() configuration.Log { // buildHTTPConfig creates HTTP server configuration from environment variables func buildHTTPConfig() (configuration.HTTP, error) { - addr := getEnvOrDefault("ATCR_HTTP_ADDR", ":5000") - debugAddr := getEnvOrDefault("ATCR_DEBUG_ADDR", ":5001") + addr := GetEnvOrDefault("ATCR_HTTP_ADDR", ":5000") + debugAddr := GetEnvOrDefault("ATCR_DEBUG_ADDR", ":5001") // HTTP secret - only needed for multipart uploads in distribution's storage driver // Since AppView is stateless and routes all storage through middleware, this isn't @@ -143,11 +143,11 @@ func buildMiddlewareConfig(defaultHoldDID string) map[string][]configuration.Mid // buildAuthConfig creates authentication configuration from environment variables func buildAuthConfig(baseURL string) (configuration.Auth, error) { // Token configuration - privateKeyPath := getEnvOrDefault("ATCR_AUTH_KEY_PATH", "/var/lib/atcr/auth/private-key.pem") - certPath := getEnvOrDefault("ATCR_AUTH_CERT_PATH", "/var/lib/atcr/auth/private-key.crt") + privateKeyPath := GetEnvOrDefault("ATCR_AUTH_KEY_PATH", "/var/lib/atcr/auth/private-key.pem") + certPath := GetEnvOrDefault("ATCR_AUTH_CERT_PATH", "/var/lib/atcr/auth/private-key.crt") // Token expiration in seconds (default: 5 minutes) - expirationStr := getEnvOrDefault("ATCR_TOKEN_EXPIRATION", "300") + expirationStr := GetEnvOrDefault("ATCR_TOKEN_EXPIRATION", "300") expiration, err := strconv.Atoi(expirationStr) if err != nil { return configuration.Auth{}, fmt.Errorf("invalid ATCR_TOKEN_EXPIRATION: %w", err) @@ -182,9 +182,9 @@ func buildHealthConfig() configuration.Health { } } -// getBaseURL determines the base URL for the service +// GetBaseURL determines the base URL for the service // Priority: ATCR_BASE_URL env var, then derived from HTTP addr -func getBaseURL(httpAddr string) string { +func GetBaseURL(httpAddr string) string { baseURL := os.Getenv("ATCR_BASE_URL") if baseURL != "" { return baseURL @@ -200,6 +200,11 @@ func getBaseURL(httpAddr string) string { return fmt.Sprintf("http://%s", httpAddr) } +// getBaseURL is the internal version used by buildAuthConfig +func getBaseURL(httpAddr string) string { + return GetBaseURL(httpAddr) +} + // getServiceName extracts service name from base URL or uses env var func getServiceName(baseURL string) string { // Check env var first @@ -224,10 +229,58 @@ func getServiceName(baseURL string) string { return "atcr.io" } -// getEnvOrDefault gets an environment variable or returns a default value -func getEnvOrDefault(key, defaultValue string) string { +// GetEnvOrDefault gets an environment variable or returns a default value +func GetEnvOrDefault(key, defaultValue string) string { if val := os.Getenv(key); val != "" { return val } return defaultValue } + +// GetStringParam extracts a string parameter from configuration.Parameters +func GetStringParam(params configuration.Parameters, key, defaultValue string) string { + if v, ok := params[key]; ok { + if s, ok := v.(string); ok { + return s + } + } + return defaultValue +} + +// GetIntParam extracts an int parameter from configuration.Parameters +func GetIntParam(params configuration.Parameters, key string, defaultValue int) int { + if v, ok := params[key]; ok { + if i, ok := v.(int); ok { + return i + } + } + return defaultValue +} + +// ExtractDefaultHoldDID extracts the default hold DID from middleware config +// Returns a DID (e.g., "did:web:hold01.atcr.io") +// To find a hold's DID, visit: https://hold-url/.well-known/did.json +func ExtractDefaultHoldDID(config *configuration.Configuration) string { + // Navigate through: middleware.registry[].options.default_hold_did + registryMiddleware, ok := config.Middleware["registry"] + if !ok { + return "" + } + + // Find atproto-resolver middleware + for _, mw := range registryMiddleware { + // Check if this is the atproto-resolver + if mw.Name != "atproto-resolver" { + continue + } + + // Extract options - options is configuration.Parameters which is map[string]any + if mw.Options != nil { + if holdDID, ok := mw.Options["default_hold_did"].(string); ok { + return holdDID + } + } + } + + return "" +} diff --git a/pkg/appview/config_test.go b/pkg/appview/config_test.go new file mode 100644 index 0000000..2eb498f --- /dev/null +++ b/pkg/appview/config_test.go @@ -0,0 +1,844 @@ +package appview + +import ( + "os" + "testing" + + "github.com/distribution/distribution/v3/configuration" +) + +func TestGetEnvOrDefault(t *testing.T) { + tests := []struct { + name string + key string + defaultValue string + envValue string + setEnv bool + want string + }{ + { + name: "env var not set", + key: "TEST_VAR_NOT_SET", + defaultValue: "default", + setEnv: false, + want: "default", + }, + { + name: "env var set to value", + key: "TEST_VAR_SET", + defaultValue: "default", + envValue: "custom", + setEnv: true, + want: "custom", + }, + { + name: "env var set to empty string", + key: "TEST_VAR_EMPTY", + defaultValue: "default", + envValue: "", + setEnv: true, + want: "default", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if tt.setEnv { + t.Setenv(tt.key, tt.envValue) + } + + got := GetEnvOrDefault(tt.key, tt.defaultValue) + if got != tt.want { + t.Errorf("GetEnvOrDefault() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestGetBaseURL(t *testing.T) { + tests := []struct { + name string + httpAddr string + envBaseURL string + setEnv bool + want string + }{ + { + name: "env var set", + httpAddr: ":5000", + envBaseURL: "https://registry.example.com", + setEnv: true, + want: "https://registry.example.com", + }, + { + name: "port only - auto detect localhost", + httpAddr: ":5000", + setEnv: false, + want: "http://127.0.0.1:5000", + }, + { + name: "full address", + httpAddr: "0.0.0.0:5000", + setEnv: false, + want: "http://0.0.0.0:5000", + }, + { + name: "custom port", + httpAddr: ":8080", + setEnv: false, + want: "http://127.0.0.1:8080", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if tt.setEnv { + t.Setenv("ATCR_BASE_URL", tt.envBaseURL) + } else { + os.Unsetenv("ATCR_BASE_URL") + } + + got := GetBaseURL(tt.httpAddr) + if got != tt.want { + t.Errorf("GetBaseURL() = %v, want %v", got, tt.want) + } + }) + } +} + +func Test_getServiceName(t *testing.T) { + tests := []struct { + name string + baseURL string + envService string + setEnv bool + want string + }{ + { + name: "env var set", + baseURL: "http://127.0.0.1:5000", + envService: "custom.registry.io", + setEnv: true, + want: "custom.registry.io", + }, + { + name: "localhost - use default", + baseURL: "http://localhost:5000", + setEnv: false, + want: "atcr.io", + }, + { + name: "127.0.0.1 - use default", + baseURL: "http://127.0.0.1:5000", + setEnv: false, + want: "atcr.io", + }, + { + name: "custom domain", + baseURL: "https://registry.example.com", + setEnv: false, + want: "registry.example.com", + }, + { + name: "domain with port", + baseURL: "https://registry.example.com:443", + setEnv: false, + want: "registry.example.com", + }, + { + name: "invalid URL - use default", + baseURL: "://invalid", + setEnv: false, + want: "atcr.io", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if tt.setEnv { + t.Setenv("ATCR_SERVICE_NAME", tt.envService) + } else { + os.Unsetenv("ATCR_SERVICE_NAME") + } + + got := getServiceName(tt.baseURL) + if got != tt.want { + t.Errorf("getServiceName() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestBuildLogConfig(t *testing.T) { + tests := []struct { + name string + envLevel string + envFormatter string + setLevel bool + setFormatter bool + wantLevel configuration.Loglevel + wantFormatter string + }{ + { + name: "defaults", + setLevel: false, + setFormatter: false, + wantLevel: "info", + wantFormatter: "text", + }, + { + name: "custom level", + envLevel: "debug", + setLevel: true, + setFormatter: false, + wantLevel: "debug", + wantFormatter: "text", + }, + { + name: "custom formatter", + envLevel: "info", + envFormatter: "json", + setLevel: true, + setFormatter: true, + wantLevel: "info", + wantFormatter: "json", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if tt.setLevel { + t.Setenv("ATCR_LOG_LEVEL", tt.envLevel) + } else { + os.Unsetenv("ATCR_LOG_LEVEL") + } + + if tt.setFormatter { + t.Setenv("ATCR_LOG_FORMATTER", tt.envFormatter) + } else { + os.Unsetenv("ATCR_LOG_FORMATTER") + } + + got := buildLogConfig() + if got.Level != tt.wantLevel { + t.Errorf("buildLogConfig().Level = %v, want %v", got.Level, tt.wantLevel) + } + if got.Formatter != tt.wantFormatter { + t.Errorf("buildLogConfig().Formatter = %v, want %v", got.Formatter, tt.wantFormatter) + } + if got.Fields["service"] != "atcr-appview" { + t.Errorf("buildLogConfig().Fields[service] = %v, want atcr-appview", got.Fields["service"]) + } + }) + } +} + +func TestBuildHTTPConfig(t *testing.T) { + tests := []struct { + name string + envAddr string + envDebugAddr string + envSecret string + setAddr bool + setDebugAddr bool + setSecret bool + wantAddr string + wantDebug string + wantSecret string // empty means "should be generated" + }{ + { + name: "defaults", + setAddr: false, + wantAddr: ":5000", + wantDebug: ":5001", + wantSecret: "", // generated + }, + { + name: "custom addr", + envAddr: ":8080", + setAddr: true, + setDebugAddr: false, + wantAddr: ":8080", + wantDebug: ":5001", + wantSecret: "", + }, + { + name: "custom debug addr", + envDebugAddr: ":9001", + setAddr: false, + setDebugAddr: true, + wantAddr: ":5000", + wantDebug: ":9001", + wantSecret: "", + }, + { + name: "custom secret", + envSecret: "my-custom-secret", + setAddr: false, + setSecret: true, + wantAddr: ":5000", + wantDebug: ":5001", + wantSecret: "my-custom-secret", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if tt.setAddr { + t.Setenv("ATCR_HTTP_ADDR", tt.envAddr) + } else { + os.Unsetenv("ATCR_HTTP_ADDR") + } + + if tt.setDebugAddr { + t.Setenv("ATCR_DEBUG_ADDR", tt.envDebugAddr) + } else { + os.Unsetenv("ATCR_DEBUG_ADDR") + } + + if tt.setSecret { + t.Setenv("REGISTRY_HTTP_SECRET", tt.envSecret) + } else { + os.Unsetenv("REGISTRY_HTTP_SECRET") + } + + got, err := buildHTTPConfig() + if err != nil { + t.Fatalf("buildHTTPConfig() error = %v", err) + } + + if got.Addr != tt.wantAddr { + t.Errorf("buildHTTPConfig().Addr = %v, want %v", got.Addr, tt.wantAddr) + } + + if got.Debug.Addr != tt.wantDebug { + t.Errorf("buildHTTPConfig().Debug.Addr = %v, want %v", got.Debug.Addr, tt.wantDebug) + } + + if tt.wantSecret == "" { + // Should be generated (64 hex chars = 32 bytes) + if len(got.Secret) != 64 { + t.Errorf("buildHTTPConfig().Secret length = %v, want 64", len(got.Secret)) + } + } else { + if got.Secret != tt.wantSecret { + t.Errorf("buildHTTPConfig().Secret = %v, want %v", got.Secret, tt.wantSecret) + } + } + + // Verify headers + if got.Headers["X-Content-Type-Options"][0] != "nosniff" { + t.Error("buildHTTPConfig() missing X-Content-Type-Options header") + } + }) + } +} + +func TestBuildStorageConfig(t *testing.T) { + got := buildStorageConfig() + + // Verify inmemory driver exists + if _, ok := got["inmemory"]; !ok { + t.Error("buildStorageConfig() missing inmemory driver") + } + + // Verify maintenance config + maintenance, ok := got["maintenance"] + if !ok { + t.Fatal("buildStorageConfig() missing maintenance config") + } + + uploadPurging, ok := maintenance["uploadpurging"] + if !ok { + t.Fatal("buildStorageConfig() missing uploadpurging config") + } + + // Verify uploadpurging is map[any]any (for distribution validation) + purging, ok := uploadPurging.(map[any]any) + if !ok { + t.Fatalf("uploadpurging is %T, want map[any]any", uploadPurging) + } + + if purging["enabled"] != false { + t.Error("uploadpurging enabled should be false") + } +} + +func TestBuildMiddlewareConfig(t *testing.T) { + tests := []struct { + name string + defaultHoldDID string + testMode bool + setTestMode bool + wantTestMode bool + }{ + { + name: "normal mode", + defaultHoldDID: "did:web:hold01.atcr.io", + setTestMode: false, + wantTestMode: false, + }, + { + name: "test mode enabled", + defaultHoldDID: "did:web:hold01.atcr.io", + testMode: true, + setTestMode: true, + wantTestMode: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if tt.setTestMode { + t.Setenv("TEST_MODE", "true") + } else { + os.Unsetenv("TEST_MODE") + } + + got := buildMiddlewareConfig(tt.defaultHoldDID) + + registryMW, ok := got["registry"] + if !ok { + t.Fatal("buildMiddlewareConfig() missing registry middleware") + } + + if len(registryMW) != 1 { + t.Fatalf("buildMiddlewareConfig() registry middleware count = %v, want 1", len(registryMW)) + } + + mw := registryMW[0] + if mw.Name != "atproto-resolver" { + t.Errorf("middleware name = %v, want atproto-resolver", mw.Name) + } + + if mw.Options["default_hold_did"] != tt.defaultHoldDID { + t.Errorf("default_hold_did = %v, want %v", mw.Options["default_hold_did"], tt.defaultHoldDID) + } + + if mw.Options["test_mode"] != tt.wantTestMode { + t.Errorf("test_mode = %v, want %v", mw.Options["test_mode"], tt.wantTestMode) + } + }) + } +} + +func TestBuildAuthConfig(t *testing.T) { + tests := []struct { + name string + baseURL string + envKeyPath string + envCertPath string + envExpiration string + setKeyPath bool + setCertPath bool + setExpiration bool + wantKeyPath string + wantCertPath string + wantExpiration int + wantRealm string + wantService string + wantError bool + }{ + { + name: "defaults", + baseURL: "http://127.0.0.1:5000", + setKeyPath: false, + setCertPath: false, + setExpiration: false, + wantKeyPath: "/var/lib/atcr/auth/private-key.pem", + wantCertPath: "/var/lib/atcr/auth/private-key.crt", + wantExpiration: 300, + wantRealm: "http://127.0.0.1:5000/auth/token", + wantService: "atcr.io", + wantError: false, + }, + { + name: "custom values", + baseURL: "https://registry.example.com", + envKeyPath: "/custom/key.pem", + envCertPath: "/custom/cert.crt", + envExpiration: "600", + setKeyPath: true, + setCertPath: true, + setExpiration: true, + wantKeyPath: "/custom/key.pem", + wantCertPath: "/custom/cert.crt", + wantExpiration: 600, + wantRealm: "https://registry.example.com/auth/token", + wantService: "registry.example.com", + wantError: false, + }, + { + name: "invalid expiration", + baseURL: "http://127.0.0.1:5000", + envExpiration: "not-a-number", + setExpiration: true, + wantError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if tt.setKeyPath { + t.Setenv("ATCR_AUTH_KEY_PATH", tt.envKeyPath) + } else { + os.Unsetenv("ATCR_AUTH_KEY_PATH") + } + + if tt.setCertPath { + t.Setenv("ATCR_AUTH_CERT_PATH", tt.envCertPath) + } else { + os.Unsetenv("ATCR_AUTH_CERT_PATH") + } + + if tt.setExpiration { + t.Setenv("ATCR_TOKEN_EXPIRATION", tt.envExpiration) + } else { + os.Unsetenv("ATCR_TOKEN_EXPIRATION") + } + + // Clear service name env var + os.Unsetenv("ATCR_SERVICE_NAME") + + got, err := buildAuthConfig(tt.baseURL) + if (err != nil) != tt.wantError { + t.Errorf("buildAuthConfig() error = %v, wantError %v", err, tt.wantError) + return + } + + if tt.wantError { + return + } + + tokenParams, ok := got["token"] + if !ok { + t.Fatal("buildAuthConfig() missing token params") + } + + if tokenParams["privatekey"] != tt.wantKeyPath { + t.Errorf("privatekey = %v, want %v", tokenParams["privatekey"], tt.wantKeyPath) + } + + if tokenParams["rootcertbundle"] != tt.wantCertPath { + t.Errorf("rootcertbundle = %v, want %v", tokenParams["rootcertbundle"], tt.wantCertPath) + } + + if tokenParams["expiration"] != tt.wantExpiration { + t.Errorf("expiration = %v, want %v", tokenParams["expiration"], tt.wantExpiration) + } + + if tokenParams["realm"] != tt.wantRealm { + t.Errorf("realm = %v, want %v", tokenParams["realm"], tt.wantRealm) + } + + if tokenParams["service"] != tt.wantService { + t.Errorf("service = %v, want %v", tokenParams["service"], tt.wantService) + } + + if tokenParams["issuer"] != tt.wantService { + t.Errorf("issuer = %v, want %v", tokenParams["issuer"], tt.wantService) + } + }) + } +} + +func TestBuildHealthConfig(t *testing.T) { + got := buildHealthConfig() + + if !got.StorageDriver.Enabled { + t.Error("buildHealthConfig().StorageDriver.Enabled = false, want true") + } + + if got.StorageDriver.Interval.Seconds() != 10 { + t.Errorf("buildHealthConfig().StorageDriver.Interval = %v, want 10s", got.StorageDriver.Interval) + } + + if got.StorageDriver.Threshold != 3 { + t.Errorf("buildHealthConfig().StorageDriver.Threshold = %v, want 3", got.StorageDriver.Threshold) + } +} + +func TestGetStringParam(t *testing.T) { + tests := []struct { + name string + params configuration.Parameters + key string + defaultValue string + want string + }{ + { + name: "string value exists", + params: configuration.Parameters{ + "foo": "bar", + }, + key: "foo", + defaultValue: "default", + want: "bar", + }, + { + name: "key does not exist", + params: configuration.Parameters{}, + key: "foo", + defaultValue: "default", + want: "default", + }, + { + name: "value is not a string", + params: configuration.Parameters{ + "foo": 123, + }, + key: "foo", + defaultValue: "default", + want: "default", + }, + { + name: "empty string value", + params: configuration.Parameters{ + "foo": "", + }, + key: "foo", + defaultValue: "default", + want: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := GetStringParam(tt.params, tt.key, tt.defaultValue) + if got != tt.want { + t.Errorf("GetStringParam() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestGetIntParam(t *testing.T) { + tests := []struct { + name string + params configuration.Parameters + key string + defaultValue int + want int + }{ + { + name: "int value exists", + params: configuration.Parameters{ + "foo": 42, + }, + key: "foo", + defaultValue: 100, + want: 42, + }, + { + name: "key does not exist", + params: configuration.Parameters{}, + key: "foo", + defaultValue: 100, + want: 100, + }, + { + name: "value is not an int", + params: configuration.Parameters{ + "foo": "not-an-int", + }, + key: "foo", + defaultValue: 100, + want: 100, + }, + { + name: "zero value", + params: configuration.Parameters{ + "foo": 0, + }, + key: "foo", + defaultValue: 100, + want: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := GetIntParam(tt.params, tt.key, tt.defaultValue) + if got != tt.want { + t.Errorf("GetIntParam() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestExtractDefaultHoldDID(t *testing.T) { + tests := []struct { + name string + config *configuration.Configuration + want string + }{ + { + name: "valid config with hold DID", + config: &configuration.Configuration{ + Middleware: map[string][]configuration.Middleware{ + "registry": { + { + Name: "atproto-resolver", + Options: configuration.Parameters{ + "default_hold_did": "did:web:hold01.atcr.io", + }, + }, + }, + }, + }, + want: "did:web:hold01.atcr.io", + }, + { + name: "no registry middleware", + config: &configuration.Configuration{ + Middleware: map[string][]configuration.Middleware{}, + }, + want: "", + }, + { + name: "no atproto-resolver middleware", + config: &configuration.Configuration{ + Middleware: map[string][]configuration.Middleware{ + "registry": { + { + Name: "other-middleware", + Options: configuration.Parameters{ + "foo": "bar", + }, + }, + }, + }, + }, + want: "", + }, + { + name: "atproto-resolver without default_hold_did", + config: &configuration.Configuration{ + Middleware: map[string][]configuration.Middleware{ + "registry": { + { + Name: "atproto-resolver", + Options: configuration.Parameters{ + "other_option": "value", + }, + }, + }, + }, + }, + want: "", + }, + { + name: "default_hold_did is not a string", + config: &configuration.Configuration{ + Middleware: map[string][]configuration.Middleware{ + "registry": { + { + Name: "atproto-resolver", + Options: configuration.Parameters{ + "default_hold_did": 123, + }, + }, + }, + }, + }, + want: "", + }, + { + name: "nil options", + config: &configuration.Configuration{ + Middleware: map[string][]configuration.Middleware{ + "registry": { + { + Name: "atproto-resolver", + Options: nil, + }, + }, + }, + }, + want: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := ExtractDefaultHoldDID(tt.config) + if got != tt.want { + t.Errorf("ExtractDefaultHoldDID() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestLoadConfigFromEnv(t *testing.T) { + tests := []struct { + name string + envHoldDID string + setHoldDID bool + wantError bool + }{ + { + name: "valid config", + envHoldDID: "did:web:hold01.atcr.io", + setHoldDID: true, + wantError: false, + }, + { + name: "missing default hold DID", + setHoldDID: false, + wantError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if tt.setHoldDID { + t.Setenv("ATCR_DEFAULT_HOLD_DID", tt.envHoldDID) + } else { + os.Unsetenv("ATCR_DEFAULT_HOLD_DID") + } + + // Clear other env vars to use defaults + os.Unsetenv("ATCR_BASE_URL") + os.Unsetenv("ATCR_SERVICE_NAME") + + got, err := LoadConfigFromEnv() + if (err != nil) != tt.wantError { + t.Errorf("LoadConfigFromEnv() error = %v, wantError %v", err, tt.wantError) + return + } + + if tt.wantError { + return + } + + // Verify config structure + if got.Version.Major() != 0 || got.Version.Minor() != 1 { + t.Errorf("version = %v, want 0.1", got.Version) + } + + if got.Log.Level != "info" { + t.Errorf("log level = %v, want info", got.Log.Level) + } + + if got.HTTP.Addr != ":5000" { + t.Errorf("HTTP addr = %v, want :5000", got.HTTP.Addr) + } + + if _, ok := got.Storage["inmemory"]; !ok { + t.Error("storage missing inmemory driver") + } + + if _, ok := got.Middleware["registry"]; !ok { + t.Error("middleware missing registry") + } + + if _, ok := got.Auth["token"]; !ok { + t.Error("auth missing token config") + } + + if !got.Health.StorageDriver.Enabled { + t.Error("health storage driver not enabled") + } + }) + } +} diff --git a/pkg/appview/db/readonly.go b/pkg/appview/db/readonly.go new file mode 100644 index 0000000..172c140 --- /dev/null +++ b/pkg/appview/db/readonly.go @@ -0,0 +1,115 @@ +package db + +import ( + "context" + "database/sql" + "fmt" + "os" + "path/filepath" + "time" + + sqlite3 "github.com/mattn/go-sqlite3" +) + +const ( + // ReadOnlyDriverName is the name of the custom SQLite driver with table authorization + ReadOnlyDriverName = "sqlite3_readonly_public" +) + +// sensitiveTables defines tables that should never be accessible from public queries +var sensitiveTables = map[string]bool{ + "oauth_sessions": true, // OAuth tokens + "ui_sessions": true, // Session IDs + "oauth_auth_requests": true, // OAuth state + "devices": true, // Device secret hashes + "pending_device_auth": true, // Pending device secrets +} + +// readOnlyAuthorizerCallback blocks access to sensitive tables +func readOnlyAuthorizerCallback(action int, arg1, arg2, dbName string) int { + // arg1 contains the table name for most operations + tableName := arg1 + + // Block any access to sensitive tables + if action == sqlite3.SQLITE_READ || action == sqlite3.SQLITE_UPDATE || + action == sqlite3.SQLITE_INSERT || action == sqlite3.SQLITE_DELETE || + action == sqlite3.SQLITE_SELECT { + if sensitiveTables[tableName] { + fmt.Printf("SECURITY: Blocked access to sensitive table '%s' (action=%d)\n", tableName, action) + return sqlite3.SQLITE_DENY + } + } + + // Allow everything else + return sqlite3.SQLITE_OK +} + +func init() { + // Register a custom SQLite driver with authorizer for read-only public queries + sql.Register(ReadOnlyDriverName, + &sqlite3.SQLiteDriver{ + ConnectHook: func(conn *sqlite3.SQLiteConn) error { + conn.RegisterAuthorizer(readOnlyAuthorizerCallback) + return nil + }, + }) +} + +// InitializeDatabase initializes the SQLite database and session store +// Returns: (read-write DB, read-only DB, session store) +func InitializeDatabase(uiEnabled bool, dbPath string) (*sql.DB, *sql.DB, *SessionStore) { + if !uiEnabled { + return nil, nil, nil + } + + // Ensure directory exists + dbDir := filepath.Dir(dbPath) + if err := os.MkdirAll(dbDir, 0700); err != nil { + fmt.Printf("Warning: Failed to create UI database directory: %v\n", err) + return nil, nil, nil + } + + // Initialize read-write database (for writes and auth operations) + database, err := InitDB(dbPath) + if err != nil { + fmt.Printf("Warning: Failed to initialize UI database: %v\n", err) + return nil, nil, nil + } + + // Open read-only connection for public queries (search, user pages, etc.) + // Uses custom driver with SQLite authorizer that blocks sensitive tables + // This prevents accidental writes and blocks access to sensitive tables even if SQL injection occurs + readOnlyDB, err := sql.Open(ReadOnlyDriverName, "file:"+dbPath+"?mode=ro") + if err != nil { + fmt.Printf("Warning: Failed to open read-only database connection: %v\n", err) + return nil, nil, nil + } + + fmt.Printf("UI database (readonly) initialized at %s\n", dbPath) + + // Create SQLite-backed session store + sessionStore := NewSessionStore(database) + + // Start cleanup goroutines for all SQLite stores + go func() { + ticker := time.NewTicker(5 * time.Minute) + defer ticker.Stop() + for range ticker.C { + ctx := context.Background() + + // Cleanup UI sessions + sessionStore.Cleanup() + + // Cleanup OAuth sessions (older than 30 days) + oauthStore := NewOAuthStore(database) + oauthStore.CleanupOldSessions(ctx, 30*24*time.Hour) + oauthStore.CleanupExpiredAuthRequests(ctx) + + // Cleanup device pending auths + deviceStore := NewDeviceStore(database) + deviceStore.CleanupExpired() + } + }() + + return database, readOnlyDB, sessionStore +} diff --git a/cmd/appview/authorizer_test.go b/pkg/appview/db/readonly_test.go similarity index 95% rename from cmd/appview/authorizer_test.go rename to pkg/appview/db/readonly_test.go index bc2a4c1..44a754a 100644 --- a/cmd/appview/authorizer_test.go +++ b/pkg/appview/db/readonly_test.go @@ -1,12 +1,10 @@ -package main +package db import ( "database/sql" "os" "path/filepath" "testing" - - "atcr.io/pkg/appview/db" ) func TestAuthorizerBlocksSensitiveTables(t *testing.T) { @@ -19,7 +17,7 @@ func TestAuthorizerBlocksSensitiveTables(t *testing.T) { defer os.Unsetenv("ATCR_UI_DATABASE_PATH") // Initialize database (creates schema) - database, err := db.InitDB(dbPath) + database, err := InitDB(dbPath) if err != nil { t.Fatalf("Failed to initialize database: %v", err) } @@ -43,7 +41,7 @@ func TestAuthorizerBlocksSensitiveTables(t *testing.T) { } // Open read-only connection with authorizer (using our custom driver) - readOnlyDB, err := sql.Open("sqlite3_readonly_public", "file:"+dbPath+"?mode=ro") + readOnlyDB, err := sql.Open(ReadOnlyDriverName, "file:"+dbPath+"?mode=ro") if err != nil { t.Fatalf("Failed to open read-only database: %v", err) } diff --git a/pkg/appview/handlers/common.go b/pkg/appview/handlers/common.go index 37254b0..2a731a7 100644 --- a/pkg/appview/handlers/common.go +++ b/pkg/appview/handlers/common.go @@ -2,6 +2,7 @@ package handlers import ( "net/http" + "strings" "atcr.io/pkg/appview/db" "atcr.io/pkg/appview/middleware" @@ -22,3 +23,11 @@ func NewPageData(r *http.Request, registryURL string) PageData { RegistryURL: registryURL, } } + +// TrimRegistryURL removes http:// or https:// prefix from a URL +// for use in Docker commands where only the host:port is needed +func TrimRegistryURL(url string) string { + url = strings.TrimPrefix(url, "https://") + url = strings.TrimPrefix(url, "http://") + return url +} diff --git a/pkg/appview/handlers/util.go b/pkg/appview/handlers/util.go deleted file mode 100644 index ca1555d..0000000 --- a/pkg/appview/handlers/util.go +++ /dev/null @@ -1,11 +0,0 @@ -package handlers - -import "strings" - -// TrimRegistryURL removes http:// or https:// prefix from a URL -// for use in Docker commands where only the host:port is needed -func TrimRegistryURL(url string) string { - url = strings.TrimPrefix(url, "https://") - url = strings.TrimPrefix(url, "http://") - return url -} diff --git a/pkg/appview/storage/proxy_blob_store.go b/pkg/appview/storage/proxy_blob_store.go index bf282c0..241f297 100644 --- a/pkg/appview/storage/proxy_blob_store.go +++ b/pkg/appview/storage/proxy_blob_store.go @@ -76,10 +76,10 @@ func resolveHoldURL(holdDID string) string { // Use HTTP for localhost/IP addresses with ports, HTTPS for domains if strings.Contains(hostname, ":") || - strings.Contains(hostname, "127.0.0.1") || - strings.Contains(hostname, "localhost") || - // Check if it's an IP address (contains only digits and dots) - (len(hostname) > 0 && (hostname[0] >= '0' && hostname[0] <= '9')) { + strings.Contains(hostname, "127.0.0.1") || + strings.Contains(hostname, "localhost") || + // Check if it's an IP address (contains only digits and dots) + (len(hostname) > 0 && (hostname[0] >= '0' && hostname[0] <= '9')) { return "http://" + hostname } return "https://" + hostname diff --git a/pkg/appview/appview.go b/pkg/appview/ui.go similarity index 100% rename from pkg/appview/appview.go rename to pkg/appview/ui.go diff --git a/pkg/appview/ui_test.go b/pkg/appview/ui_test.go new file mode 100644 index 0000000..75fc125 --- /dev/null +++ b/pkg/appview/ui_test.go @@ -0,0 +1,601 @@ +package appview + +import ( + "bytes" + "strings" + "testing" + "time" +) + +func TestTimeAgo(t *testing.T) { + now := time.Now() + + tests := []struct { + name string + time time.Time + expected string + }{ + { + name: "just now - 30 seconds ago", + time: now.Add(-30 * time.Second), + expected: "just now", + }, + { + name: "1 minute ago", + time: now.Add(-1 * time.Minute), + expected: "1 minute ago", + }, + { + name: "5 minutes ago", + time: now.Add(-5 * time.Minute), + expected: "5 minutes ago", + }, + { + name: "45 minutes ago", + time: now.Add(-45 * time.Minute), + expected: "45 minutes ago", + }, + { + name: "1 hour ago", + time: now.Add(-1 * time.Hour), + expected: "1 hour ago", + }, + { + name: "3 hours ago", + time: now.Add(-3 * time.Hour), + expected: "3 hours ago", + }, + { + name: "23 hours ago", + time: now.Add(-23 * time.Hour), + expected: "23 hours ago", + }, + { + name: "1 day ago", + time: now.Add(-24 * time.Hour), + expected: "1 day ago", + }, + { + name: "5 days ago", + time: now.Add(-5 * 24 * time.Hour), + expected: "5 days ago", + }, + { + name: "30 days ago", + time: now.Add(-30 * 24 * time.Hour), + expected: "30 days ago", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // Get fresh template for each test case + tmpl, err := Templates() + if err != nil { + t.Fatalf("Templates() error = %v", err) + } + + // Execute template using timeAgo function + templateStr := `{{ timeAgo . }}` + buf := new(bytes.Buffer) + temp, err := tmpl.New("test").Parse(templateStr) + if err != nil { + t.Fatalf("Failed to parse template: %v", err) + } + + err = temp.Execute(buf, tt.time) + if err != nil { + t.Fatalf("Failed to execute template: %v", err) + } + + got := buf.String() + if got != tt.expected { + t.Errorf("timeAgo() = %q, want %q", got, tt.expected) + } + }) + } +} + +func TestHumanizeBytes(t *testing.T) { + tests := []struct { + name string + bytes int64 + expected string + }{ + { + name: "0 bytes", + bytes: 0, + expected: "0 B", + }, + { + name: "512 bytes", + bytes: 512, + expected: "512 B", + }, + { + name: "1023 bytes", + bytes: 1023, + expected: "1023 B", + }, + { + name: "1 KB", + bytes: 1024, + expected: "1.0 KB", + }, + { + name: "1.5 KB", + bytes: 1536, + expected: "1.5 KB", + }, + { + name: "1 MB", + bytes: 1024 * 1024, + expected: "1.0 MB", + }, + { + name: "2.5 MB", + bytes: 2621440, // 2.5 * 1024 * 1024 + expected: "2.5 MB", + }, + { + name: "1 GB", + bytes: 1024 * 1024 * 1024, + expected: "1.0 GB", + }, + { + name: "5.2 GB", + bytes: 5583457485, // ~5.2 GB + expected: "5.2 GB", + }, + { + name: "1 TB", + bytes: 1024 * 1024 * 1024 * 1024, + expected: "1.0 TB", + }, + { + name: "1.5 PB", + bytes: 1688849860263936, // 1.5 PB + expected: "1.5 PB", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // Get fresh template for each test case + tmpl, err := Templates() + if err != nil { + t.Fatalf("Templates() error = %v", err) + } + + templateStr := `{{ humanizeBytes . }}` + buf := new(bytes.Buffer) + temp, err := tmpl.New("test").Parse(templateStr) + if err != nil { + t.Fatalf("Failed to parse template: %v", err) + } + + err = temp.Execute(buf, tt.bytes) + if err != nil { + t.Fatalf("Failed to execute template: %v", err) + } + + got := buf.String() + if got != tt.expected { + t.Errorf("humanizeBytes(%d) = %q, want %q", tt.bytes, got, tt.expected) + } + }) + } +} + +func TestTruncateDigest(t *testing.T) { + tests := []struct { + name string + digest string + length int + expected string + }{ + { + name: "short digest - no truncation needed", + digest: "sha256:abc", + length: 20, + expected: "sha256:abc", + }, + { + name: "truncate to 12 chars", + digest: "sha256:abcdef123456789", + length: 12, + expected: "sha256:abcde...", + }, + { + name: "truncate to 8 chars", + digest: "sha256:1234567890abcdef", + length: 8, + expected: "sha256:1...", + }, + { + name: "exact length match", + digest: "sha256:abc", + length: 10, + expected: "sha256:abc", + }, + { + name: "empty digest", + digest: "", + length: 10, + expected: "", + }, + { + name: "long sha256 digest", + digest: "sha256:f1c8f6a4b7e9d2c0a3f5b8e1d4c7a0b3e6f9c2d5a8b1e4f7c0d3a6b9e2f5c8a1", + length: 16, + expected: "sha256:f1c8f6a4b...", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // Get fresh template for each test case + tmpl, err := Templates() + if err != nil { + t.Fatalf("Templates() error = %v", err) + } + + templateStr := `{{ truncateDigest .Digest .Length }}` + buf := new(bytes.Buffer) + temp, err := tmpl.New("test").Parse(templateStr) + if err != nil { + t.Fatalf("Failed to parse template: %v", err) + } + + data := struct { + Digest string + Length int + }{ + Digest: tt.digest, + Length: tt.length, + } + + err = temp.Execute(buf, data) + if err != nil { + t.Fatalf("Failed to execute template: %v", err) + } + + got := buf.String() + if got != tt.expected { + t.Errorf("truncateDigest(%q, %d) = %q, want %q", tt.digest, tt.length, got, tt.expected) + } + }) + } +} + +func TestFirstChar(t *testing.T) { + tests := []struct { + name string + input string + expected string + }{ + { + name: "normal string", + input: "hello", + expected: "h", + }, + { + name: "uppercase", + input: "World", + expected: "W", + }, + { + name: "single character", + input: "a", + expected: "a", + }, + { + name: "empty string", + input: "", + expected: "?", + }, + { + name: "unicode character", + input: "😀 emoji", + expected: "😀", + }, + { + name: "chinese character", + input: "你好", + expected: "你", + }, + { + name: "number", + input: "123", + expected: "1", + }, + { + name: "special character", + input: "@user", + expected: "@", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // Get fresh template for each test case + tmpl, err := Templates() + if err != nil { + t.Fatalf("Templates() error = %v", err) + } + + templateStr := `{{ firstChar . }}` + buf := new(bytes.Buffer) + temp, err := tmpl.New("test").Parse(templateStr) + if err != nil { + t.Fatalf("Failed to parse template: %v", err) + } + + err = temp.Execute(buf, tt.input) + if err != nil { + t.Fatalf("Failed to execute template: %v", err) + } + + got := buf.String() + if got != tt.expected { + t.Errorf("firstChar(%q) = %q, want %q", tt.input, got, tt.expected) + } + }) + } +} + +func TestTrimPrefix(t *testing.T) { + tests := []struct { + name string + prefix string + input string + expected string + }{ + { + name: "trim sha256 prefix", + prefix: "sha256:", + input: "sha256:abcdef123456", + expected: "abcdef123456", + }, + { + name: "no prefix match", + prefix: "sha256:", + input: "md5:abcdef123456", + expected: "md5:abcdef123456", + }, + { + name: "empty prefix", + prefix: "", + input: "hello", + expected: "hello", + }, + { + name: "empty string", + prefix: "prefix:", + input: "", + expected: "", + }, + { + name: "prefix longer than string", + prefix: "very-long-prefix", + input: "short", + expected: "short", + }, + { + name: "exact match", + prefix: "prefix", + input: "prefix", + expected: "", + }, + { + name: "partial prefix match", + prefix: "sha256:", + input: "sha25", + expected: "sha25", + }, + { + name: "trim docker.io prefix", + prefix: "docker.io/", + input: "docker.io/library/alpine", + expected: "library/alpine", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // Get fresh template for each test case + tmpl, err := Templates() + if err != nil { + t.Fatalf("Templates() error = %v", err) + } + + templateStr := `{{ trimPrefix .Prefix .Input }}` + buf := new(bytes.Buffer) + temp, err := tmpl.New("test").Parse(templateStr) + if err != nil { + t.Fatalf("Failed to parse template: %v", err) + } + + data := struct { + Prefix string + Input string + }{ + Prefix: tt.prefix, + Input: tt.input, + } + + err = temp.Execute(buf, data) + if err != nil { + t.Fatalf("Failed to execute template: %v", err) + } + + got := buf.String() + if got != tt.expected { + t.Errorf("trimPrefix(%q, %q) = %q, want %q", tt.prefix, tt.input, got, tt.expected) + } + }) + } +} + +func TestTemplates(t *testing.T) { + tmpl, err := Templates() + if err != nil { + t.Fatalf("Templates() error = %v", err) + } + + if tmpl == nil { + t.Fatal("Templates() returned nil template") + } + + // Test that all expected templates are loaded + expectedTemplates := []string{ + "base.html", + "nav", + "repo-card", + "repository", + "home.html", + "search.html", + "user.html", + "login.html", + "settings.html", + "install.html", + "manifest-modal", + "push-list.html", + } + + for _, name := range expectedTemplates { + t.Run("template_"+name, func(t *testing.T) { + temp := tmpl.Lookup(name) + if temp == nil { + t.Errorf("Expected template %q not found", name) + } + }) + } +} + +func TestTemplateExecution_RepoCard(t *testing.T) { + tmpl, err := Templates() + if err != nil { + t.Fatalf("Templates() error = %v", err) + } + + // Sample data for repo-card template + data := struct { + OwnerHandle string + Repository string + IconURL string + Description string + StarCount int + PullCount int + }{ + OwnerHandle: "alice.bsky.social", + Repository: "myapp", + IconURL: "", + Description: "A cool container image", + StarCount: 42, + PullCount: 1337, + } + + buf := new(bytes.Buffer) + err = tmpl.ExecuteTemplate(buf, "repo-card", data) + if err != nil { + t.Fatalf("Failed to execute repo-card template: %v", err) + } + + output := buf.String() + + // Verify expected content in output + expectedContent := []string{ + "alice.bsky.social", + "myapp", + "A cool container image", + "42", // star count + "1337", // pull count + "featured-icon-placeholder", // no icon URL provided + } + + for _, expected := range expectedContent { + if !strings.Contains(output, expected) { + t.Errorf("Template output missing expected content %q", expected) + } + } + + // Verify firstChar function is working + if !strings.Contains(output, ">m<") { // first char of "myapp" + t.Error("Template output missing firstChar result") + } +} + +func TestTemplateExecution_WithFuncMap(t *testing.T) { + // Test that templates can use FuncMap functions + tests := []struct { + name string + templateStr string + data interface{} + expectInOutput string + }{ + { + name: "timeAgo in template", + templateStr: `{{ define "test1" }}{{ timeAgo . }}{{ end }}`, + data: time.Now().Add(-5 * time.Minute), + expectInOutput: "5 minutes ago", + }, + { + name: "humanizeBytes in template", + templateStr: `{{ define "test2" }}{{ humanizeBytes . }}{{ end }}`, + data: int64(1024 * 1024 * 10), // 10 MB + expectInOutput: "10.0 MB", + }, + { + name: "multiple functions in template", + templateStr: `{{ define "test3" }}{{ truncateDigest .Digest 12 }} - {{ firstChar .Name }}{{ end }}`, + data: struct { + Digest string + Name string + }{ + Digest: "sha256:abcdef1234567890", + Name: "myapp", + }, + expectInOutput: "sha256:abcde... - m", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // Get fresh template for each test case + tmpl, err := Templates() + if err != nil { + t.Fatalf("Templates() error = %v", err) + } + + temp, err := tmpl.Parse(tt.templateStr) + if err != nil { + t.Fatalf("Failed to parse template: %v", err) + } + + buf := new(bytes.Buffer) + // Extract the template name from the define + templateName := strings.Split(strings.TrimPrefix(tt.templateStr, `{{ define "`), `"`)[0] + err = temp.ExecuteTemplate(buf, templateName, tt.data) + if err != nil { + t.Fatalf("Failed to execute template: %v", err) + } + + output := buf.String() + if !strings.Contains(output, tt.expectInOutput) { + t.Errorf("Template output %q does not contain expected %q", output, tt.expectInOutput) + } + }) + } +} + +func TestStaticHandler(t *testing.T) { + handler := StaticHandler() + if handler == nil { + t.Fatal("StaticHandler() returned nil") + } + + // Test that it returns an http.Handler + // Further testing would require HTTP request/response testing + // which is typically done in integration tests +} diff --git a/pkg/atproto/lexicon.go b/pkg/atproto/lexicon.go index 4d1becf..0dc88b4 100644 --- a/pkg/atproto/lexicon.go +++ b/pkg/atproto/lexicon.go @@ -1,7 +1,5 @@ package atproto -//go:generate go run github.com/whyrusleeping/cbor-gen --map-encoding CrewRecord CaptainRecord - import ( "encoding/base64" "encoding/json" @@ -221,59 +219,6 @@ func NewHoldRecord(endpoint, owner string, public bool) *HoldRecord { } } -// HoldCrewRecord represents membership in a storage hold -// Stored in the hold owner's PDS (not the crew member's PDS) to ensure owner maintains full control -// Owner can add/remove crew members by creating/deleting these records in their own PDS -// Supports both explicit DIDs (with backlinks) and pattern-based matching (wildcards, handle globs) -type HoldCrewRecord struct { - // Type should be "io.atcr.hold.crew" - Type string `json:"$type"` - - // Hold is the AT URI of the hold record - // e.g., "at://did:plc:owner/io.atcr.hold/hold1" - Hold string `json:"hold"` - - // Member is the DID of the crew member (optional, for explicit access) - // Exactly one of Member or MemberPattern must be set - Member *string `json:"member,omitempty"` - - // MemberPattern is a pattern for matching multiple users (optional, for pattern-based access) - // Supports wildcards: "*" (all users), "*.domain.com" (handle glob) - // Exactly one of Member or MemberPattern must be set - MemberPattern *string `json:"memberPattern,omitempty"` - - // Role defines permissions: "owner", "write", "read" - Role string `json:"role"` - - // ExpiresAt is optional expiration for this membership - ExpiresAt *time.Time `json:"expiresAt,omitempty"` - - // AddedAt timestamp - AddedAt time.Time `json:"createdAt"` -} - -// NewHoldCrewRecord creates a new hold crew record with explicit DID -func NewHoldCrewRecord(hold, member, role string) *HoldCrewRecord { - return &HoldCrewRecord{ - Type: HoldCrewCollection, - Hold: hold, - Member: &member, - Role: role, - AddedAt: time.Now(), - } -} - -// NewHoldCrewRecordWithPattern creates a new hold crew record with pattern matching -func NewHoldCrewRecordWithPattern(hold, pattern, role string) *HoldCrewRecord { - return &HoldCrewRecord{ - Type: HoldCrewCollection, - Hold: hold, - MemberPattern: &pattern, - Role: role, - AddedAt: time.Now(), - } -} - // SailorProfileRecord represents a user's profile with registry preferences // Stored in the user's PDS to configure default hold and other settings type SailorProfileRecord struct { @@ -390,6 +335,11 @@ func ResolveHoldDIDFromURL(holdURL string) string { return "did:web:" + hostname } +// isDID checks if a string is a DID (starts with "did:") +func isDID(s string) bool { + return len(s) > 4 && s[:4] == "did:" +} + // ============================================================================= // Embedded PDS Types (Hold Service) // ============================================================================= diff --git a/pkg/atproto/lexicon_test.go b/pkg/atproto/lexicon_test.go new file mode 100644 index 0000000..7611ae5 --- /dev/null +++ b/pkg/atproto/lexicon_test.go @@ -0,0 +1,683 @@ +package atproto + +import ( + "encoding/json" + "strings" + "testing" + "time" +) + +func TestNewManifestRecord(t *testing.T) { + validOCIManifest := `{ + "schemaVersion": 2, + "mediaType": "application/vnd.oci.image.manifest.v1+json", + "config": { + "mediaType": "application/vnd.oci.image.config.v1+json", + "digest": "sha256:config123", + "size": 1234 + }, + "layers": [ + { + "mediaType": "application/vnd.oci.image.layer.v1.tar+gzip", + "digest": "sha256:layer1", + "size": 5678 + }, + { + "mediaType": "application/vnd.oci.image.layer.v1.tar+gzip", + "digest": "sha256:layer2", + "size": 9012 + } + ], + "annotations": { + "org.opencontainers.image.created": "2025-01-01T00:00:00Z" + } + }` + + manifestWithSubject := `{ + "schemaVersion": 2, + "mediaType": "application/vnd.oci.image.manifest.v1+json", + "config": { + "mediaType": "application/vnd.oci.image.config.v1+json", + "digest": "sha256:config123", + "size": 1234 + }, + "layers": [], + "subject": { + "mediaType": "application/vnd.oci.image.manifest.v1+json", + "digest": "sha256:subject123", + "size": 4321 + } + }` + + tests := []struct { + name string + repository string + digest string + ociManifest string + wantErr bool + checkFunc func(*testing.T, *ManifestRecord) + }{ + { + name: "valid OCI manifest", + repository: "myapp", + digest: "sha256:abc123", + ociManifest: validOCIManifest, + wantErr: false, + checkFunc: func(t *testing.T, record *ManifestRecord) { + if record.Type != ManifestCollection { + t.Errorf("Type = %v, want %v", record.Type, ManifestCollection) + } + if record.Repository != "myapp" { + t.Errorf("Repository = %v, want myapp", record.Repository) + } + if record.Digest != "sha256:abc123" { + t.Errorf("Digest = %v, want sha256:abc123", record.Digest) + } + if record.SchemaVersion != 2 { + t.Errorf("SchemaVersion = %v, want 2", record.SchemaVersion) + } + if record.MediaType != "application/vnd.oci.image.manifest.v1+json" { + t.Errorf("MediaType = %v, want application/vnd.oci.image.manifest.v1+json", record.MediaType) + } + if record.Config.Digest != "sha256:config123" { + t.Errorf("Config.Digest = %v, want sha256:config123", record.Config.Digest) + } + if record.Config.Size != 1234 { + t.Errorf("Config.Size = %v, want 1234", record.Config.Size) + } + if len(record.Layers) != 2 { + t.Fatalf("len(Layers) = %v, want 2", len(record.Layers)) + } + if record.Layers[0].Digest != "sha256:layer1" { + t.Errorf("Layers[0].Digest = %v, want sha256:layer1", record.Layers[0].Digest) + } + if record.Layers[1].Digest != "sha256:layer2" { + t.Errorf("Layers[1].Digest = %v, want sha256:layer2", record.Layers[1].Digest) + } + if record.Annotations["org.opencontainers.image.created"] != "2025-01-01T00:00:00Z" { + t.Errorf("Annotations missing expected key") + } + if record.CreatedAt.IsZero() { + t.Error("CreatedAt should not be zero") + } + if record.Subject != nil { + t.Error("Subject should be nil") + } + }, + }, + { + name: "manifest with subject", + repository: "myapp", + digest: "sha256:abc123", + ociManifest: manifestWithSubject, + wantErr: false, + checkFunc: func(t *testing.T, record *ManifestRecord) { + if record.Subject == nil { + t.Fatal("Subject should not be nil") + } + if record.Subject.Digest != "sha256:subject123" { + t.Errorf("Subject.Digest = %v, want sha256:subject123", record.Subject.Digest) + } + if record.Subject.Size != 4321 { + t.Errorf("Subject.Size = %v, want 4321", record.Subject.Size) + } + }, + }, + { + name: "invalid JSON", + repository: "myapp", + digest: "sha256:abc123", + ociManifest: "not valid json", + wantErr: true, + }, + { + name: "invalid config JSON", + repository: "myapp", + digest: "sha256:abc123", + ociManifest: `{"schemaVersion": 2, "mediaType": "test", "config": "not-an-object", "layers": []}`, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := NewManifestRecord(tt.repository, tt.digest, []byte(tt.ociManifest)) + if (err != nil) != tt.wantErr { + t.Errorf("NewManifestRecord() error = %v, wantErr %v", err, tt.wantErr) + return + } + + if !tt.wantErr && tt.checkFunc != nil { + tt.checkFunc(t, got) + } + }) + } +} + +func TestNewTagRecord(t *testing.T) { + before := time.Now() + record := NewTagRecord("myapp", "latest", "sha256:abc123") + after := time.Now() + + if record.Type != TagCollection { + t.Errorf("Type = %v, want %v", record.Type, TagCollection) + } + + if record.Repository != "myapp" { + t.Errorf("Repository = %v, want myapp", record.Repository) + } + + if record.Tag != "latest" { + t.Errorf("Tag = %v, want latest", record.Tag) + } + + if record.ManifestDigest != "sha256:abc123" { + t.Errorf("ManifestDigest = %v, want sha256:abc123", record.ManifestDigest) + } + + if record.UpdatedAt.Before(before) || record.UpdatedAt.After(after) { + t.Errorf("UpdatedAt = %v, want between %v and %v", record.UpdatedAt, before, after) + } +} + +func TestNewHoldRecord(t *testing.T) { + tests := []struct { + name string + endpoint string + owner string + public bool + }{ + { + name: "public hold", + endpoint: "https://hold1.example.com", + owner: "did:plc:alice123", + public: true, + }, + { + name: "private hold", + endpoint: "https://hold2.example.com", + owner: "did:plc:bob456", + public: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + before := time.Now() + record := NewHoldRecord(tt.endpoint, tt.owner, tt.public) + after := time.Now() + + if record.Type != HoldCollection { + t.Errorf("Type = %v, want %v", record.Type, HoldCollection) + } + + if record.Endpoint != tt.endpoint { + t.Errorf("Endpoint = %v, want %v", record.Endpoint, tt.endpoint) + } + + if record.Owner != tt.owner { + t.Errorf("Owner = %v, want %v", record.Owner, tt.owner) + } + + if record.Public != tt.public { + t.Errorf("Public = %v, want %v", record.Public, tt.public) + } + + if record.CreatedAt.Before(before) || record.CreatedAt.After(after) { + t.Errorf("CreatedAt = %v, want between %v and %v", record.CreatedAt, before, after) + } + }) + } +} + +func TestNewSailorProfileRecord(t *testing.T) { + tests := []struct { + name string + defaultHold string + }{ + { + name: "with default hold DID", + defaultHold: "did:web:hold01.atcr.io", + }, + { + name: "with default hold URL", + defaultHold: "https://hold01.atcr.io", + }, + { + name: "empty default hold", + defaultHold: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + before := time.Now() + record := NewSailorProfileRecord(tt.defaultHold) + after := time.Now() + + if record.Type != SailorProfileCollection { + t.Errorf("Type = %v, want %v", record.Type, SailorProfileCollection) + } + + if record.DefaultHold != tt.defaultHold { + t.Errorf("DefaultHold = %v, want %v", record.DefaultHold, tt.defaultHold) + } + + if record.CreatedAt.Before(before) || record.CreatedAt.After(after) { + t.Errorf("CreatedAt = %v, want between %v and %v", record.CreatedAt, before, after) + } + + if record.UpdatedAt.Before(before) || record.UpdatedAt.After(after) { + t.Errorf("UpdatedAt = %v, want between %v and %v", record.UpdatedAt, before, after) + } + + // CreatedAt and UpdatedAt should be equal for new records + if !record.CreatedAt.Equal(record.UpdatedAt) { + t.Errorf("CreatedAt (%v) != UpdatedAt (%v)", record.CreatedAt, record.UpdatedAt) + } + }) + } +} + +func TestNewStarRecord(t *testing.T) { + before := time.Now() + record := NewStarRecord("did:plc:alice123", "myapp") + after := time.Now() + + if record.Type != StarCollection { + t.Errorf("Type = %v, want %v", record.Type, StarCollection) + } + + if record.Subject.DID != "did:plc:alice123" { + t.Errorf("Subject.DID = %v, want did:plc:alice123", record.Subject.DID) + } + + if record.Subject.Repository != "myapp" { + t.Errorf("Subject.Repository = %v, want myapp", record.Subject.Repository) + } + + if record.CreatedAt.Before(before) || record.CreatedAt.After(after) { + t.Errorf("CreatedAt = %v, want between %v and %v", record.CreatedAt, before, after) + } +} + +func TestStarRecordKey(t *testing.T) { + tests := []struct { + name string + ownerDID string + repository string + wantPrefix string // Expected prefix for validation + }{ + { + name: "simple key", + ownerDID: "did:plc:alice123", + repository: "myapp", + }, + { + name: "long DID and repo", + ownerDID: "did:plc:abcdefghijklmnopqrstuvwxyz123456", + repository: "my-very-long-repository-name", + }, + { + name: "special characters in repo", + ownerDID: "did:plc:alice123", + repository: "my-app_test.v1", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + key := StarRecordKey(tt.ownerDID, tt.repository) + + // Key should be non-empty + if key == "" { + t.Error("StarRecordKey() returned empty string") + } + + // Key should be base64 URL-encoded (no padding) + if strings.Contains(key, "=") { + t.Errorf("StarRecordKey() = %v, should not contain padding", key) + } + + // Should be deterministic + key2 := StarRecordKey(tt.ownerDID, tt.repository) + if key != key2 { + t.Errorf("StarRecordKey() not deterministic: %v != %v", key, key2) + } + + // Should be different for different inputs + differentKey := StarRecordKey(tt.ownerDID, tt.repository+"different") + if key == differentKey { + t.Error("StarRecordKey() should be different for different inputs") + } + }) + } +} + +func TestParseStarRecordKey(t *testing.T) { + tests := []struct { + name string + ownerDID string + repository string + wantErr bool + }{ + { + name: "valid key", + ownerDID: "did:plc:alice123", + repository: "myapp", + wantErr: false, + }, + { + name: "key with special characters", + ownerDID: "did:plc:alice123", + repository: "my-app_test.v1", + wantErr: false, + }, + { + name: "long values", + ownerDID: "did:plc:abcdefghijklmnopqrstuvwxyz123456", + repository: "my-very-long-repository-name", + wantErr: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // Generate key + key := StarRecordKey(tt.ownerDID, tt.repository) + + // Parse it back + gotDID, gotRepo, err := ParseStarRecordKey(key) + if (err != nil) != tt.wantErr { + t.Errorf("ParseStarRecordKey() error = %v, wantErr %v", err, tt.wantErr) + return + } + + if !tt.wantErr { + if gotDID != tt.ownerDID { + t.Errorf("ParseStarRecordKey() DID = %v, want %v", gotDID, tt.ownerDID) + } + if gotRepo != tt.repository { + t.Errorf("ParseStarRecordKey() repository = %v, want %v", gotRepo, tt.repository) + } + } + }) + } +} + +func TestParseStarRecordKey_Invalid(t *testing.T) { + tests := []struct { + name string + rkey string + }{ + { + name: "invalid base64", + rkey: "not!!!valid!!!base64", + }, + { + name: "no separator - base64 encoded text without slash", + rkey: "bm9zZXBhcmF0b3I", // base64 of "noseparator" (no "/" in the decoded value) + }, + { + name: "empty string", + rkey: "", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, _, err := ParseStarRecordKey(tt.rkey) + if err == nil { + t.Error("ParseStarRecordKey() expected error for invalid input") + } + }) + } +} + +func TestResolveHoldDIDFromURL(t *testing.T) { + tests := []struct { + name string + holdURL string + want string + }{ + { + name: "https URL", + holdURL: "https://hold01.atcr.io", + want: "did:web:hold01.atcr.io", + }, + { + name: "http URL", + holdURL: "http://hold01.atcr.io", + want: "did:web:hold01.atcr.io", + }, + { + name: "URL with trailing slash", + holdURL: "https://hold01.atcr.io/", + want: "did:web:hold01.atcr.io", + }, + { + name: "URL with path", + holdURL: "https://hold01.atcr.io/some/path", + want: "did:web:hold01.atcr.io", + }, + { + name: "URL with port", + holdURL: "https://hold01.atcr.io:8080", + want: "did:web:hold01.atcr.io:8080", + }, + { + name: "already a did:web", + holdURL: "did:web:hold01.atcr.io", + want: "did:web:hold01.atcr.io", + }, + { + name: "already a did:plc", + holdURL: "did:plc:abc123", + want: "did:plc:abc123", + }, + { + name: "empty string", + holdURL: "", + want: "", + }, + { + name: "localhost", + holdURL: "http://localhost:8080", + want: "did:web:localhost:8080", + }, + { + name: "IP address", + holdURL: "http://192.168.1.1:8080", + want: "did:web:192.168.1.1:8080", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := ResolveHoldDIDFromURL(tt.holdURL) + if got != tt.want { + t.Errorf("ResolveHoldDIDFromURL() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestIsDID(t *testing.T) { + tests := []struct { + name string + s string + want bool + }{ + { + name: "valid did:web", + s: "did:web:example.com", + want: true, + }, + { + name: "valid did:plc", + s: "did:plc:abc123", + want: true, + }, + { + name: "valid did:key", + s: "did:key:z6Mkfriq", + want: true, + }, + { + name: "not a DID - URL", + s: "https://example.com", + want: false, + }, + { + name: "not a DID - short string", + s: "did", + want: false, + }, + { + name: "not a DID - empty", + s: "", + want: false, + }, + { + name: "not a DID - almost", + s: "did:", + want: false, + }, + { + name: "not a DID - plain text", + s: "hello world", + want: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := isDID(tt.s) + if got != tt.want { + t.Errorf("isDID() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestManifestRecord_JSONSerialization(t *testing.T) { + // Create a manifest record + ociManifest := `{ + "schemaVersion": 2, + "mediaType": "application/vnd.oci.image.manifest.v1+json", + "config": { + "mediaType": "application/vnd.oci.image.config.v1+json", + "digest": "sha256:config123", + "size": 1234 + }, + "layers": [ + { + "mediaType": "application/vnd.oci.image.layer.v1.tar+gzip", + "digest": "sha256:layer1", + "size": 5678 + } + ] + }` + + record, err := NewManifestRecord("myapp", "sha256:abc123", []byte(ociManifest)) + if err != nil { + t.Fatalf("NewManifestRecord() error = %v", err) + } + + // Add hold DID + record.HoldDID = "did:web:hold01.atcr.io" + + // Serialize to JSON + jsonData, err := json.Marshal(record) + if err != nil { + t.Fatalf("json.Marshal() error = %v", err) + } + + // Deserialize from JSON + var decoded ManifestRecord + if err := json.Unmarshal(jsonData, &decoded); err != nil { + t.Fatalf("json.Unmarshal() error = %v", err) + } + + // Verify fields + if decoded.Type != record.Type { + t.Errorf("Type = %v, want %v", decoded.Type, record.Type) + } + if decoded.Repository != record.Repository { + t.Errorf("Repository = %v, want %v", decoded.Repository, record.Repository) + } + if decoded.Digest != record.Digest { + t.Errorf("Digest = %v, want %v", decoded.Digest, record.Digest) + } + if decoded.HoldDID != record.HoldDID { + t.Errorf("HoldDID = %v, want %v", decoded.HoldDID, record.HoldDID) + } + if decoded.Config.Digest != record.Config.Digest { + t.Errorf("Config.Digest = %v, want %v", decoded.Config.Digest, record.Config.Digest) + } + if len(decoded.Layers) != len(record.Layers) { + t.Errorf("len(Layers) = %v, want %v", len(decoded.Layers), len(record.Layers)) + } +} + +func TestBlobReference_JSONSerialization(t *testing.T) { + blob := BlobReference{ + MediaType: "application/vnd.oci.image.layer.v1.tar+gzip", + Digest: "sha256:abc123", + Size: 12345, + URLs: []string{"https://s3.example.com/blob"}, + Annotations: map[string]string{ + "key": "value", + }, + } + + // Serialize + jsonData, err := json.Marshal(blob) + if err != nil { + t.Fatalf("json.Marshal() error = %v", err) + } + + // Deserialize + var decoded BlobReference + if err := json.Unmarshal(jsonData, &decoded); err != nil { + t.Fatalf("json.Unmarshal() error = %v", err) + } + + // Verify + if decoded.MediaType != blob.MediaType { + t.Errorf("MediaType = %v, want %v", decoded.MediaType, blob.MediaType) + } + if decoded.Digest != blob.Digest { + t.Errorf("Digest = %v, want %v", decoded.Digest, blob.Digest) + } + if decoded.Size != blob.Size { + t.Errorf("Size = %v, want %v", decoded.Size, blob.Size) + } +} + +func TestStarSubject_JSONSerialization(t *testing.T) { + subject := StarSubject{ + DID: "did:plc:alice123", + Repository: "myapp", + } + + // Serialize + jsonData, err := json.Marshal(subject) + if err != nil { + t.Fatalf("json.Marshal() error = %v", err) + } + + // Deserialize + var decoded StarSubject + if err := json.Unmarshal(jsonData, &decoded); err != nil { + t.Fatalf("json.Unmarshal() error = %v", err) + } + + // Verify + if decoded.DID != subject.DID { + t.Errorf("DID = %v, want %v", decoded.DID, subject.DID) + } + if decoded.Repository != subject.Repository { + t.Errorf("Repository = %v, want %v", decoded.Repository, subject.Repository) + } +} diff --git a/pkg/atproto/manifest_store.go b/pkg/atproto/manifest_store.go index f4f048e..c4fb003 100644 --- a/pkg/atproto/manifest_store.go +++ b/pkg/atproto/manifest_store.go @@ -21,14 +21,14 @@ type DatabaseMetrics interface { // ManifestStore implements distribution.ManifestService // It stores manifests in ATProto as records type ManifestStore struct { - client *Client - repository string - holdEndpoint string // Hold service endpoint URL (for legacy, to be deprecated) - holdDID string // Hold service DID (primary reference) - did string // User's DID for cache key - lastFetchedHoldDID string // Hold DID from most recently fetched manifest (for pull) - blobStore distribution.BlobStore // Blob store for fetching config during push - database DatabaseMetrics // Database for metrics tracking + client *Client + repository string + holdEndpoint string // Hold service endpoint URL (for legacy, to be deprecated) + holdDID string // Hold service DID (primary reference) + did string // User's DID for cache key + lastFetchedHoldDID string // Hold DID from most recently fetched manifest (for pull) + blobStore distribution.BlobStore // Blob store for fetching config during push + database DatabaseMetrics // Database for metrics tracking } // NewManifestStore creates a new ATProto-backed manifest store diff --git a/pkg/atproto/profile.go b/pkg/atproto/profile.go index dcf5db2..cfa18b7 100644 --- a/pkg/atproto/profile.go +++ b/pkg/atproto/profile.go @@ -55,7 +55,7 @@ func GetProfile(ctx context.Context, client *Client) (*SailorProfileRecord, erro record, err := client.GetRecord(ctx, SailorProfileCollection, ProfileRKey) if err != nil { // Check if it's a 404 (profile doesn't exist) - if isNotFoundError(err) { + if errors.Is(err, ErrRecordNotFound) { return nil, nil } return nil, fmt.Errorf("failed to get profile: %w", err) @@ -104,11 +104,6 @@ func GetProfile(ctx context.Context, client *Client) (*SailorProfileRecord, erro return &profile, nil } -// isDID checks if a string is a DID (starts with "did:") -func isDID(s string) bool { - return len(s) > 4 && s[:4] == "did:" -} - // UpdateProfile updates the user's profile // Normalizes defaultHold to DID format before saving func UpdateProfile(ctx context.Context, client *Client, profile *SailorProfileRecord) error { @@ -125,8 +120,3 @@ func UpdateProfile(ctx context.Context, client *Client, profile *SailorProfileRe } return nil } - -// isNotFoundError checks if an error is a record not found error -func isNotFoundError(err error) bool { - return errors.Is(err, ErrRecordNotFound) -} diff --git a/pkg/hold/blobstore_adapter.go b/pkg/hold/blobstore_adapter.go index 09a894f..16ca1bf 100644 --- a/pkg/hold/blobstore_adapter.go +++ b/pkg/hold/blobstore_adapter.go @@ -105,8 +105,8 @@ func (b *HoldServiceBlobStore) GetPartUploadURL(ctx context.Context, uploadID st URL: fmt.Sprintf("%s/xrpc/com.atproto.repo.uploadBlob", b.service.config.Server.PublicURL), Method: "PUT", Headers: map[string]string{ - "X-Upload-Id": uploadID, - "X-Part-Number": fmt.Sprintf("%d", partNumber), + "X-Upload-Id": uploadID, + "X-Part-Number": fmt.Sprintf("%d", partNumber), }, }, nil } diff --git a/pkg/hold/multipart.go b/pkg/hold/multipart.go index 1cedddc..86eb655 100644 --- a/pkg/hold/multipart.go +++ b/pkg/hold/multipart.go @@ -333,4 +333,3 @@ func (s *HoldService) AbortMultipartUploadWithManager(ctx context.Context, sessi log.Printf("Aborted buffered multipart: uploadID=%s", session.UploadID) return nil } -