mirror of
https://tangled.org/evan.jarrett.net/at-container-registry
synced 2026-09-03 08:46:57 +00:00
335 lines
9.8 KiB
Go
335 lines
9.8 KiB
Go
package middleware
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"testing"
|
|
|
|
"github.com/distribution/distribution/v3"
|
|
"github.com/distribution/reference"
|
|
"github.com/stretchr/testify/assert"
|
|
"github.com/stretchr/testify/require"
|
|
|
|
"atcr.io/pkg/atproto"
|
|
)
|
|
|
|
// mockNamespace is a mock implementation of distribution.Namespace
|
|
type mockNamespace struct {
|
|
distribution.Namespace
|
|
repositories map[string]distribution.Repository
|
|
}
|
|
|
|
func (m *mockNamespace) Repository(ctx context.Context, name reference.Named) (distribution.Repository, error) {
|
|
if m.repositories == nil {
|
|
return nil, fmt.Errorf("repository not found: %s", name.Name())
|
|
}
|
|
if repo, ok := m.repositories[name.Name()]; ok {
|
|
return repo, nil
|
|
}
|
|
return nil, fmt.Errorf("repository not found: %s", name.Name())
|
|
}
|
|
|
|
func (m *mockNamespace) Repositories(ctx context.Context, repos []string, last string) (int, error) {
|
|
// Return empty result for mock
|
|
return 0, nil
|
|
}
|
|
|
|
func (m *mockNamespace) Blobs() distribution.BlobEnumerator {
|
|
return nil
|
|
}
|
|
|
|
func (m *mockNamespace) BlobStatter() distribution.BlobStatter {
|
|
return nil
|
|
}
|
|
|
|
func TestSetGlobalRefresher(t *testing.T) {
|
|
// Test that SetGlobalRefresher doesn't panic
|
|
SetGlobalRefresher(nil)
|
|
// If we get here without panic, test passes
|
|
}
|
|
|
|
func TestSetGlobalDatabase(t *testing.T) {
|
|
SetGlobalDatabase(nil)
|
|
// If we get here without panic, test passes
|
|
}
|
|
|
|
func TestSetGlobalAuthorizer(t *testing.T) {
|
|
SetGlobalAuthorizer(nil)
|
|
// If we get here without panic, test passes
|
|
}
|
|
|
|
// TestInitATProtoResolver tests the initialization function
|
|
func TestInitATProtoResolver(t *testing.T) {
|
|
ctx := context.Background()
|
|
mockNS := &mockNamespace{}
|
|
|
|
tests := []struct {
|
|
name string
|
|
options map[string]any
|
|
wantErr bool
|
|
}{
|
|
{
|
|
name: "with default hold DID",
|
|
options: map[string]any{
|
|
"default_hold_did": "did:web:hold01.atcr.io",
|
|
"base_url": "https://atcr.io",
|
|
"test_mode": false,
|
|
},
|
|
wantErr: false,
|
|
},
|
|
{
|
|
name: "with test mode enabled",
|
|
options: map[string]any{
|
|
"default_hold_did": "did:web:hold01.atcr.io",
|
|
"base_url": "https://atcr.io",
|
|
"test_mode": true,
|
|
},
|
|
wantErr: false,
|
|
},
|
|
{
|
|
name: "without options",
|
|
options: map[string]any{},
|
|
wantErr: false,
|
|
},
|
|
}
|
|
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
ns, err := initATProtoResolver(ctx, mockNS, nil, tt.options)
|
|
if tt.wantErr {
|
|
assert.Error(t, err)
|
|
return
|
|
}
|
|
|
|
require.NoError(t, err)
|
|
assert.NotNil(t, ns)
|
|
|
|
resolver, ok := ns.(*NamespaceResolver)
|
|
require.True(t, ok, "expected NamespaceResolver type")
|
|
|
|
if holdDID, ok := tt.options["default_hold_did"].(string); ok {
|
|
assert.Equal(t, holdDID, resolver.defaultHoldDID)
|
|
}
|
|
if baseURL, ok := tt.options["base_url"].(string); ok {
|
|
assert.Equal(t, baseURL, resolver.baseURL)
|
|
}
|
|
if testMode, ok := tt.options["test_mode"].(bool); ok {
|
|
assert.Equal(t, testMode, resolver.testMode)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
// TestAuthErrorMessage tests the error message formatting
|
|
func TestAuthErrorMessage(t *testing.T) {
|
|
resolver := &NamespaceResolver{
|
|
baseURL: "https://atcr.io",
|
|
}
|
|
|
|
err := resolver.authErrorMessage("OAuth session expired")
|
|
assert.Contains(t, err.Error(), "OAuth session expired")
|
|
assert.Contains(t, err.Error(), "https://atcr.io/auth/oauth/login")
|
|
}
|
|
|
|
// TestFindHoldDID_DefaultFallback tests default hold DID fallback
|
|
func TestFindHoldDID_DefaultFallback(t *testing.T) {
|
|
// Start a mock PDS server that returns 404 for profile and empty list for holds
|
|
mockPDS := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
if r.URL.Path == "/xrpc/com.atproto.repo.getRecord" {
|
|
// Profile not found
|
|
w.WriteHeader(http.StatusNotFound)
|
|
return
|
|
}
|
|
if r.URL.Path == "/xrpc/com.atproto.repo.listRecords" {
|
|
// Empty hold records
|
|
w.Header().Set("Content-Type", "application/json")
|
|
json.NewEncoder(w).Encode(map[string]any{
|
|
"records": []any{},
|
|
})
|
|
return
|
|
}
|
|
w.WriteHeader(http.StatusNotFound)
|
|
}))
|
|
defer mockPDS.Close()
|
|
|
|
resolver := &NamespaceResolver{
|
|
defaultHoldDID: "did:web:default.atcr.io",
|
|
}
|
|
|
|
ctx := context.Background()
|
|
holdDID, _ := resolver.findHoldDIDAndProfile(ctx, "did:plc:test123", mockPDS.URL)
|
|
|
|
assert.Equal(t, "did:web:default.atcr.io", holdDID, "should fall back to default hold DID")
|
|
}
|
|
|
|
// TestFindHoldDID_SailorProfile tests hold discovery from sailor profile
|
|
func TestFindHoldDID_SailorProfile(t *testing.T) {
|
|
// Start a mock PDS server that returns a sailor profile
|
|
mockPDS := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
if r.URL.Path == "/xrpc/com.atproto.repo.getRecord" {
|
|
// Return sailor profile with defaultHold
|
|
profile := atproto.NewSailorProfileRecord("did:web:user.hold.io")
|
|
w.Header().Set("Content-Type", "application/json")
|
|
json.NewEncoder(w).Encode(map[string]any{
|
|
"value": profile,
|
|
})
|
|
return
|
|
}
|
|
w.WriteHeader(http.StatusNotFound)
|
|
}))
|
|
defer mockPDS.Close()
|
|
|
|
resolver := &NamespaceResolver{
|
|
defaultHoldDID: "did:web:default.atcr.io",
|
|
testMode: false,
|
|
}
|
|
|
|
ctx := context.Background()
|
|
holdDID, _ := resolver.findHoldDIDAndProfile(ctx, "did:plc:test123", mockPDS.URL)
|
|
|
|
assert.Equal(t, "did:web:user.hold.io", holdDID, "should use sailor profile's defaultHold")
|
|
}
|
|
|
|
// TestFindHoldDID_Priority tests the priority order
|
|
func TestFindHoldDID_Priority(t *testing.T) {
|
|
// Start a mock PDS server that returns both profile and hold records
|
|
mockPDS := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
if r.URL.Path == "/xrpc/com.atproto.repo.getRecord" {
|
|
// Return sailor profile with defaultHold (highest priority)
|
|
profile := atproto.NewSailorProfileRecord("did:web:profile.hold.io")
|
|
w.Header().Set("Content-Type", "application/json")
|
|
json.NewEncoder(w).Encode(map[string]any{
|
|
"value": profile,
|
|
})
|
|
return
|
|
}
|
|
w.WriteHeader(http.StatusNotFound)
|
|
}))
|
|
defer mockPDS.Close()
|
|
|
|
resolver := &NamespaceResolver{
|
|
defaultHoldDID: "did:web:default.atcr.io",
|
|
}
|
|
|
|
ctx := context.Background()
|
|
holdDID, _ := resolver.findHoldDIDAndProfile(ctx, "did:plc:test123", mockPDS.URL)
|
|
|
|
// Profile should take priority over hold records and default
|
|
assert.Equal(t, "did:web:profile.hold.io", holdDID, "should prioritize sailor profile over hold records")
|
|
}
|
|
|
|
// TestFindHoldDID_TestModeFallback tests test mode fallback when hold unreachable
|
|
func TestFindHoldDID_TestModeFallback(t *testing.T) {
|
|
// Start a mock PDS server that returns a profile with unreachable hold
|
|
mockPDS := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
if r.URL.Path == "/xrpc/com.atproto.repo.getRecord" {
|
|
// Return sailor profile with an unreachable hold
|
|
profile := atproto.NewSailorProfileRecord("did:web:unreachable.hold.io")
|
|
w.Header().Set("Content-Type", "application/json")
|
|
json.NewEncoder(w).Encode(map[string]any{
|
|
"value": profile,
|
|
})
|
|
return
|
|
}
|
|
w.WriteHeader(http.StatusNotFound)
|
|
}))
|
|
defer mockPDS.Close()
|
|
|
|
resolver := &NamespaceResolver{
|
|
defaultHoldDID: "did:web:default.atcr.io",
|
|
testMode: true, // Test mode enabled
|
|
}
|
|
|
|
ctx := context.Background()
|
|
holdDID, _ := resolver.findHoldDIDAndProfile(ctx, "did:plc:test123", mockPDS.URL)
|
|
|
|
// In test mode with unreachable hold, should fall back to default
|
|
assert.Equal(t, "did:web:default.atcr.io", holdDID, "should fall back to default in test mode when hold unreachable")
|
|
}
|
|
|
|
// TestIsHoldReachable tests the hold reachability check
|
|
func TestIsHoldReachable(t *testing.T) {
|
|
// Mock hold server with DID document
|
|
mockHold := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
if r.URL.Path == "/.well-known/did.json" {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
json.NewEncoder(w).Encode(map[string]any{
|
|
"id": "did:web:reachable.hold.io",
|
|
})
|
|
return
|
|
}
|
|
w.WriteHeader(http.StatusNotFound)
|
|
}))
|
|
defer mockHold.Close()
|
|
|
|
resolver := &NamespaceResolver{}
|
|
|
|
ctx := context.Background()
|
|
|
|
t.Run("reachable hold", func(t *testing.T) {
|
|
// Use URL format directly — DID resolution requires real identity directory
|
|
reachable := resolver.isHoldReachable(ctx, mockHold.URL)
|
|
assert.True(t, reachable, "should detect reachable hold")
|
|
})
|
|
|
|
t.Run("unreachable hold", func(t *testing.T) {
|
|
reachable := resolver.isHoldReachable(ctx, "did:web:nonexistent.example.com")
|
|
assert.False(t, reachable, "should detect unreachable hold")
|
|
})
|
|
}
|
|
|
|
// TestRepositoryCaching tests that repositories are cached by DID+name
|
|
func TestRepositoryCaching(t *testing.T) {
|
|
// This test requires integration with actual repository resolution
|
|
// For now, we test that the cache key format is correct
|
|
did := "did:plc:test123"
|
|
repoName := "myapp"
|
|
expectedKey := "did:plc:test123:myapp"
|
|
|
|
cacheKey := did + ":" + repoName
|
|
assert.Equal(t, expectedKey, cacheKey, "cache key should be DID:reponame")
|
|
}
|
|
|
|
// TestNamespaceResolver_Repositories tests delegation to underlying namespace
|
|
func TestNamespaceResolver_Repositories(t *testing.T) {
|
|
mockNS := &mockNamespace{}
|
|
resolver := &NamespaceResolver{
|
|
Namespace: mockNS,
|
|
}
|
|
|
|
ctx := context.Background()
|
|
repos := []string{}
|
|
|
|
// Test delegation (mockNamespace doesn't implement this, so it will return 0, nil)
|
|
n, err := resolver.Repositories(ctx, repos, "")
|
|
assert.NoError(t, err)
|
|
assert.Equal(t, 0, n)
|
|
}
|
|
|
|
// TestNamespaceResolver_Blobs tests delegation to underlying namespace
|
|
func TestNamespaceResolver_Blobs(t *testing.T) {
|
|
mockNS := &mockNamespace{}
|
|
resolver := &NamespaceResolver{
|
|
Namespace: mockNS,
|
|
}
|
|
|
|
// Should not panic
|
|
blobs := resolver.Blobs()
|
|
assert.Nil(t, blobs, "mockNamespace returns nil")
|
|
}
|
|
|
|
// TestNamespaceResolver_BlobStatter tests delegation to underlying namespace
|
|
func TestNamespaceResolver_BlobStatter(t *testing.T) {
|
|
mockNS := &mockNamespace{}
|
|
resolver := &NamespaceResolver{
|
|
Namespace: mockNS,
|
|
}
|
|
|
|
// Should not panic
|
|
statter := resolver.BlobStatter()
|
|
assert.Nil(t, statter, "mockNamespace returns nil")
|
|
}
|