post to bluesky when manifests uploaded. linting fixes

This commit is contained in:
Evan Jarrett
2025-10-23 12:24:04 -05:00
parent 220022c9c5
commit 751fa1a3f0
41 changed files with 2701 additions and 490 deletions
+7 -1
View File
@@ -417,7 +417,13 @@ func (s *DeviceStore) CleanupExpiredContext(ctx context.Context) error {
func generateUserCode() string {
chars := "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
code := make([]byte, 8)
rand.Read(code)
if _, err := rand.Read(code); err != nil {
// Fallback to timestamp-based generation if crypto rand fails
now := time.Now().UnixNano()
for i := range code {
code[i] = byte(now >> (i * 8))
}
}
for i := range code {
code[i] = chars[int(code[i])%len(chars)]
}
+3 -3
View File
@@ -85,7 +85,7 @@ func TestInvalidateSessionsWithMismatchedScopes(t *testing.T) {
}
// Verify mismatched session was deleted
retrieved, err = store.GetSession(ctx, mismatchedSession.AccountDID, mismatchedSession.SessionID)
_, err = store.GetSession(ctx, mismatchedSession.AccountDID, mismatchedSession.SessionID)
if err == nil {
t.Error("Expected session to be deleted (should error), but got no error")
}
@@ -154,7 +154,7 @@ func TestInvalidateSessionsWithMismatchedScopes(t *testing.T) {
}
// Verify malformed session was deleted
retrieved, err = store.GetSession(ctx, parsedDID, "malformed")
_, err = store.GetSession(ctx, parsedDID, "malformed")
if err == nil {
t.Error("Expected malformed session to be deleted, but got no error")
}
@@ -284,7 +284,7 @@ func TestOAuthStoreSessionLifecycle(t *testing.T) {
}
// Verify deletion
retrieved, err = store.GetSession(ctx, did, "test_session_id")
_, err = store.GetSession(ctx, did, "test_session_id")
if err == nil {
t.Error("Expected error after deletion, got nil")
}
+3 -1
View File
@@ -13,7 +13,9 @@ func TestAuthorizerBlocksSensitiveTables(t *testing.T) {
dbPath := filepath.Join(tmpDir, "test.db")
// Set environment for database path
os.Setenv("ATCR_UI_DATABASE_PATH", dbPath)
if err := os.Setenv("ATCR_UI_DATABASE_PATH", dbPath); err != nil {
t.Fatalf("Failed to set environment variable: %v", err)
}
defer os.Unsetenv("ATCR_UI_DATABASE_PATH")
// Initialize database (creates schema)
+1 -1
View File
@@ -402,7 +402,7 @@ func (b *BackfillWorker) reconcileAnnotations(ctx context.Context, did string, p
}
// Update annotations from newest manifest only
if manifestRecord.Annotations != nil && len(manifestRecord.Annotations) > 0 {
if len(manifestRecord.Annotations) > 0 {
// Filter out empty annotations
hasData := false
for _, value := range manifestRecord.Annotations {
+7 -2
View File
@@ -26,6 +26,9 @@ import (
"atcr.io/pkg/auth/token"
)
// holdDIDKey is the context key for storing hold DID
const holdDIDKey contextKey = "hold.did"
// Global variables for initialization only
// These are set by main.go during startup and copied into NamespaceResolver instances.
// After initialization, request handling uses the NamespaceResolver's instance fields.
@@ -131,12 +134,13 @@ func (nr *NamespaceResolver) Repository(ctx context.Context, name reference.Name
}
did := ident.DID.String()
handle := ident.Handle.String()
pdsEndpoint := ident.PDSEndpoint()
if pdsEndpoint == "" {
return nil, fmt.Errorf("no PDS endpoint found for %s", identityStr)
}
fmt.Printf("DEBUG [registry/middleware]: Resolved identity: did=%s, pds=%s, handle=%s\n", did, pdsEndpoint, ident.Handle.String())
fmt.Printf("DEBUG [registry/middleware]: Resolved identity: did=%s, pds=%s, handle=%s\n", did, pdsEndpoint, handle)
// Query for hold DID - either user's hold or default hold service
holdDID := nr.findHoldDID(ctx, did, pdsEndpoint)
@@ -144,7 +148,7 @@ func (nr *NamespaceResolver) Repository(ctx context.Context, name reference.Name
// This is a fatal configuration error - registry cannot function without a hold service
return nil, fmt.Errorf("no hold DID configured: ensure default_hold_did is set in middleware config")
}
ctx = context.WithValue(ctx, "hold.did", holdDID)
ctx = context.WithValue(ctx, holdDIDKey, holdDID)
// Get service token for hold authentication
// Check cache first to avoid unnecessary PDS calls on every request
@@ -308,6 +312,7 @@ func (nr *NamespaceResolver) Repository(ctx context.Context, name reference.Name
// Bundle all context into a single RegistryContext struct
registryCtx := &storage.RegistryContext{
DID: did,
Handle: handle,
HoldDID: holdDID,
PDSEndpoint: pdsEndpoint,
Repository: repositoryName,
+1
View File
@@ -17,6 +17,7 @@ type DatabaseMetrics interface {
type RegistryContext struct {
// Per-request identity and routing information
DID string // User's DID (e.g., "did:plc:abc123")
Handle string // User's handle (e.g., "alice.bsky.social")
HoldDID string // Hold service DID (e.g., "did:web:hold01.atcr.io")
PDSEndpoint string // User's PDS endpoint URL
Repository string // Image repository name (e.g., "debian")
@@ -1,56 +1,51 @@
package atproto
package storage
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"maps"
"net/http"
"strings"
"atcr.io/pkg/atproto"
"github.com/distribution/distribution/v3"
"github.com/opencontainers/go-digest"
)
// DatabaseMetrics interface for tracking push and pull counts
type DatabaseMetrics interface {
IncrementPushCount(did, repository string) error
IncrementPullCount(did, repository string) error
// HoldNotifier interface for notifying holds about manifest uploads
type HoldNotifier interface {
GetServiceToken(ctx context.Context, userDID, audienceDID string) (string, error)
}
// 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
ctx *RegistryContext // Context with user/hold info
notifier HoldNotifier // OAuth refresher for getting service tokens
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
func NewManifestStore(client *Client, repository string, holdEndpoint string, holdDID string, did string, blobStore distribution.BlobStore, database DatabaseMetrics) *ManifestStore {
func NewManifestStore(ctx *RegistryContext, notifier HoldNotifier, blobStore distribution.BlobStore) *ManifestStore {
return &ManifestStore{
client: client,
repository: repository,
holdEndpoint: holdEndpoint,
holdDID: holdDID,
did: did,
blobStore: blobStore,
database: database,
ctx: ctx,
notifier: notifier,
blobStore: blobStore,
}
}
// Exists checks if a manifest exists by digest
func (s *ManifestStore) Exists(ctx context.Context, dgst digest.Digest) (bool, error) {
rkey := digestToRKey(dgst)
_, err := s.client.GetRecord(ctx, ManifestCollection, rkey)
_, err := s.ctx.ATProtoClient.GetRecord(ctx, atproto.ManifestCollection, rkey)
if err != nil {
// If not found, return false without error
if errors.Is(err, ErrRecordNotFound) {
if errors.Is(err, atproto.ErrRecordNotFound) {
return false, nil
}
return false, err
@@ -61,15 +56,15 @@ func (s *ManifestStore) Exists(ctx context.Context, dgst digest.Digest) (bool, e
// Get retrieves a manifest by digest
func (s *ManifestStore) Get(ctx context.Context, dgst digest.Digest, options ...distribution.ManifestServiceOption) (distribution.Manifest, error) {
rkey := digestToRKey(dgst)
record, err := s.client.GetRecord(ctx, ManifestCollection, rkey)
record, err := s.ctx.ATProtoClient.GetRecord(ctx, atproto.ManifestCollection, rkey)
if err != nil {
return nil, distribution.ErrManifestUnknownRevision{
Name: s.repository,
Name: s.ctx.Repository,
Revision: dgst,
}
}
var manifestRecord ManifestRecord
var manifestRecord atproto.ManifestRecord
if err := json.Unmarshal(record.Value, &manifestRecord); err != nil {
return nil, fmt.Errorf("failed to unmarshal manifest record: %w", err)
}
@@ -82,24 +77,24 @@ func (s *ManifestStore) Get(ctx context.Context, dgst digest.Digest, options ...
s.lastFetchedHoldDID = manifestRecord.HoldDID
} else if manifestRecord.HoldEndpoint != "" {
// Legacy format: URL reference - convert to DID
s.lastFetchedHoldDID = ResolveHoldDIDFromURL(manifestRecord.HoldEndpoint)
s.lastFetchedHoldDID = atproto.ResolveHoldDIDFromURL(manifestRecord.HoldEndpoint)
}
var ociManifest []byte
// New records: Download blob from ATProto blob storage
if manifestRecord.ManifestBlob != nil && manifestRecord.ManifestBlob.Ref.Link != "" {
ociManifest, err = s.client.GetBlob(ctx, manifestRecord.ManifestBlob.Ref.Link)
ociManifest, err = s.ctx.ATProtoClient.GetBlob(ctx, manifestRecord.ManifestBlob.Ref.Link)
if err != nil {
return nil, fmt.Errorf("failed to download manifest blob: %w", err)
}
}
// Track pull count (increment asynchronously to avoid blocking the response)
if s.database != nil {
if s.ctx.Database != nil {
go func() {
if err := s.database.IncrementPullCount(s.did, s.repository); err != nil {
fmt.Printf("WARNING: Failed to increment pull count for %s/%s: %v\n", s.did, s.repository, err)
if err := s.ctx.Database.IncrementPullCount(s.ctx.DID, s.ctx.Repository); err != nil {
fmt.Printf("WARNING: Failed to increment pull count for %s/%s: %v\n", s.ctx.DID, s.ctx.Repository, err)
}
}()
}
@@ -125,21 +120,25 @@ func (s *ManifestStore) Put(ctx context.Context, manifest distribution.Manifest,
dgst := digest.FromBytes(payload)
// Upload manifest as blob to PDS
blobRef, err := s.client.UploadBlob(ctx, payload, mediaType)
blobRef, err := s.ctx.ATProtoClient.UploadBlob(ctx, payload, mediaType)
if err != nil {
return "", fmt.Errorf("failed to upload manifest blob: %w", err)
}
// Create manifest record with structured metadata
manifestRecord, err := NewManifestRecord(s.repository, dgst.String(), payload)
manifestRecord, err := atproto.NewManifestRecord(s.ctx.Repository, dgst.String(), payload)
if err != nil {
return "", fmt.Errorf("failed to create manifest record: %w", err)
}
// Set the blob reference, hold DID, and hold endpoint
manifestRecord.ManifestBlob = blobRef
manifestRecord.HoldDID = s.holdDID // Primary reference (DID)
manifestRecord.HoldEndpoint = s.holdEndpoint // Legacy reference (URL) for backward compat
manifestRecord.HoldDID = s.ctx.HoldDID // Primary reference (DID)
// Resolve hold endpoint from DID for backward compatibility
if holdEndpoint, err := resolveDIDToHTTPSEndpoint(s.ctx.HoldDID); err == nil {
manifestRecord.HoldEndpoint = holdEndpoint // Legacy reference (URL) for backward compat
}
// Extract Dockerfile labels from config blob and add to annotations
// Only for image manifests (not manifest lists which don't have config blobs)
@@ -166,40 +165,51 @@ func (s *ManifestStore) Put(ctx context.Context, manifest distribution.Manifest,
// Store manifest record in ATProto
rkey := digestToRKey(dgst)
_, err = s.client.PutRecord(ctx, ManifestCollection, rkey, manifestRecord)
_, err = s.ctx.ATProtoClient.PutRecord(ctx, atproto.ManifestCollection, rkey, manifestRecord)
if err != nil {
return "", fmt.Errorf("failed to store manifest record in ATProto: %w", err)
}
// Track push count (increment asynchronously to avoid blocking the response)
if s.database != nil {
if s.ctx.Database != nil {
go func() {
if err := s.database.IncrementPushCount(s.did, s.repository); err != nil {
fmt.Printf("WARNING: Failed to increment push count for %s/%s: %v\n", s.did, s.repository, err)
if err := s.ctx.Database.IncrementPushCount(s.ctx.DID, s.ctx.Repository); err != nil {
fmt.Printf("WARNING: Failed to increment push count for %s/%s: %v\n", s.ctx.DID, s.ctx.Repository, err)
}
}()
}
// Also handle tag if specified
var tag string
for _, option := range options {
if tagOpt, ok := option.(distribution.WithTagOption); ok {
tag := tagOpt.Tag
tagRecord := NewTagRecord(s.client.DID(), s.repository, tag, dgst.String())
tagRKey := RepositoryTagToRKey(s.repository, tag)
_, err = s.client.PutRecord(ctx, TagCollection, tagRKey, tagRecord)
tag = tagOpt.Tag
tagRecord := atproto.NewTagRecord(s.ctx.ATProtoClient.DID(), s.ctx.Repository, tag, dgst.String())
tagRKey := atproto.RepositoryTagToRKey(s.ctx.Repository, tag)
_, err = s.ctx.ATProtoClient.PutRecord(ctx, atproto.TagCollection, tagRKey, tagRecord)
if err != nil {
return "", fmt.Errorf("failed to store tag in ATProto: %w", err)
}
}
}
// Notify hold about manifest upload (for layer tracking and Bluesky posts)
// Do this asynchronously to avoid blocking the push
if tag != "" && s.notifier != nil && s.ctx.Handle != "" {
go func() {
if err := s.notifyHoldAboutManifest(context.Background(), manifestRecord, tag, dgst.String()); err != nil {
fmt.Printf("WARNING: Failed to notify hold about manifest: %v\n", err)
}
}()
}
return dgst, nil
}
// Delete removes a manifest
func (s *ManifestStore) Delete(ctx context.Context, dgst digest.Digest) error {
rkey := digestToRKey(dgst)
return s.client.DeleteRecord(ctx, ManifestCollection, rkey)
return s.ctx.ATProtoClient.DeleteRecord(ctx, atproto.ManifestCollection, rkey)
}
// digestToRKey converts a digest to an ATProto record key
@@ -209,40 +219,6 @@ func digestToRKey(dgst digest.Digest) string {
return dgst.Encoded()
}
// RepositoryTagToRKey converts a repository and tag to an ATProto record key
// ATProto record keys must match: ^[a-zA-Z0-9._~-]{1,512}$
func RepositoryTagToRKey(repository, tag string) string {
// Combine repository and tag to create a unique key
// Replace invalid characters: slashes become tildes (~)
// We use tilde instead of dash to avoid ambiguity with repository names that contain hyphens
key := fmt.Sprintf("%s_%s", repository, tag)
// Replace / with ~ (slash not allowed in rkeys, tilde is allowed and unlikely in repo names)
key = strings.ReplaceAll(key, "/", "~")
return key
}
// RKeyToRepositoryTag converts an ATProto record key back to repository and tag
// This is the inverse of RepositoryTagToRKey
// Note: If the tag contains underscores, this will split on the LAST underscore
func RKeyToRepositoryTag(rkey string) (repository, tag string) {
// Find the last underscore to split repository and tag
lastUnderscore := strings.LastIndex(rkey, "_")
if lastUnderscore == -1 {
// No underscore found - treat entire string as tag with empty repository
return "", rkey
}
repository = rkey[:lastUnderscore]
tag = rkey[lastUnderscore+1:]
// Convert tildes back to slashes in repository (tilde was used to encode slashes)
repository = strings.ReplaceAll(repository, "~", "/")
return repository, tag
}
// GetLastFetchedHoldDID returns the hold DID from the most recently fetched manifest
// This is used by the routing repository to cache the hold for blob requests
func (s *ManifestStore) GetLastFetchedHoldDID() string {
@@ -291,3 +267,106 @@ func (s *ManifestStore) extractConfigLabels(ctx context.Context, configDigestStr
return configJSON.Config.Labels, nil
}
// resolveDIDToHTTPSEndpoint resolves a DID to an HTTPS endpoint
// Currently supports did:web only (e.g., did:web:hold01.atcr.io → https://hold01.atcr.io)
func resolveDIDToHTTPSEndpoint(did string) (string, error) {
if !strings.HasPrefix(did, "did:web:") {
return "", fmt.Errorf("only did:web is supported, got: %s", did)
}
// Extract hostname from did:web
hostname := strings.TrimPrefix(did, "did:web:")
// Handle port notation (did:web:example.com:8080 → https://example.com:8080)
hostname = strings.ReplaceAll(hostname, ":", ":")
return "https://" + hostname, nil
}
// notifyHoldAboutManifest notifies the hold service about a manifest upload
// This enables the hold to create layer records and Bluesky posts
func (s *ManifestStore) notifyHoldAboutManifest(ctx context.Context, manifestRecord *atproto.ManifestRecord, tag, manifestDigest string) error {
// Skip if no notifier configured
if s.notifier == nil {
return nil
}
// Resolve hold DID to HTTP endpoint
// For did:web, this is straightforward (e.g., did:web:hold01.atcr.io → https://hold01.atcr.io)
holdEndpoint, err := resolveDIDToHTTPSEndpoint(s.ctx.HoldDID)
if err != nil {
return fmt.Errorf("failed to resolve hold DID %s: %w", s.ctx.HoldDID, err)
}
// Get service token from user's PDS for hold authentication
serviceToken, err := s.notifier.GetServiceToken(ctx, s.ctx.DID, s.ctx.HoldDID)
if err != nil {
return fmt.Errorf("failed to get service token: %w", err)
}
// Build notification request
notifyReq := map[string]any{
"repository": s.ctx.Repository,
"tag": tag,
"userDid": s.ctx.DID,
"userHandle": s.ctx.Handle,
"manifest": map[string]any{
"mediaType": manifestRecord.MediaType,
"config": map[string]any{
"digest": manifestRecord.Config.Digest,
"size": manifestRecord.Config.Size,
},
"layers": func() []map[string]any {
layers := make([]map[string]any, len(manifestRecord.Layers))
for i, layer := range manifestRecord.Layers {
layers[i] = map[string]any{
"digest": layer.Digest,
"size": layer.Size,
"mediaType": layer.MediaType,
}
}
return layers
}(),
},
}
// Marshal request
reqBody, err := json.Marshal(notifyReq)
if err != nil {
return fmt.Errorf("failed to marshal notification request: %w", err)
}
// Send notification to hold
req, err := http.NewRequestWithContext(
ctx,
"POST",
holdEndpoint+atproto.HoldNotifyManifest,
bytes.NewReader(reqBody),
)
if err != nil {
return fmt.Errorf("failed to create HTTP request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+serviceToken)
resp, err := http.DefaultClient.Do(req)
if err != nil {
return fmt.Errorf("failed to send notification: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
body, _ := io.ReadAll(resp.Body)
return fmt.Errorf("hold notification failed: status %d, body: %s", resp.StatusCode, body)
}
// Parse response (optional logging)
var notifyResp map[string]any
if err := json.NewDecoder(resp.Body).Decode(&notifyResp); err == nil {
fmt.Printf("INFO: Hold notification successful for %s:%s - %+v\n", s.ctx.Repository, tag, notifyResp)
}
return nil
}
@@ -1,4 +1,4 @@
package atproto
package storage
import (
"context"
@@ -7,6 +7,7 @@ import (
"net/http"
"testing"
"atcr.io/pkg/atproto"
"github.com/distribution/distribution/v3"
"github.com/opencontainers/go-digest"
)
@@ -92,16 +93,15 @@ func (m *mockBlobStore) Open(ctx context.Context, dgst digest.Digest) (io.ReadSe
return nil, nil // Not needed for current tests
}
// mockATProtoClient mocks the ATProto client for testing
type mockATProtoClient struct {
records map[string]map[string]interface{} // collection -> rkey -> record
blobs map[string][]byte // cid -> blob data
}
func newMockATProtoClient() *mockATProtoClient {
return &mockATProtoClient{
records: make(map[string]map[string]interface{}),
blobs: make(map[string][]byte),
// mockRegistryContext creates a mock RegistryContext for testing
func mockRegistryContext(client *atproto.Client, repository, holdDID, did, handle string, database DatabaseMetrics) *RegistryContext {
return &RegistryContext{
ATProtoClient: client,
Repository: repository,
HoldDID: holdDID,
DID: did,
Handle: handle,
Database: database,
}
}
@@ -134,159 +134,26 @@ func TestDigestToRKey(t *testing.T) {
}
}
// TestRepositoryTagToRKey tests repository+tag to record key conversion
func TestRepositoryTagToRKey(t *testing.T) {
tests := []struct {
name string
repository string
tag string
want string
}{
{
name: "simple repo and tag",
repository: "myapp",
tag: "latest",
want: "myapp_latest",
},
{
name: "repo with namespace",
repository: "org/myapp",
tag: "v1.0.0",
want: "org~myapp_v1.0.0",
},
{
name: "tag with underscore",
repository: "myapp",
tag: "test_tag",
want: "myapp_test_tag",
},
{
name: "deep namespace",
repository: "a/b/c/myapp",
tag: "prod",
want: "a~b~c~myapp_prod",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := RepositoryTagToRKey(tt.repository, tt.tag)
if got != tt.want {
t.Errorf("RepositoryTagToRKey() = %v, want %v", got, tt.want)
}
})
}
}
// TestRKeyToRepositoryTag tests converting record key back to repository and tag
func TestRKeyToRepositoryTag(t *testing.T) {
tests := []struct {
name string
rkey string
wantRepository string
wantTag string
}{
{
name: "simple key",
rkey: "myapp_latest",
wantRepository: "myapp",
wantTag: "latest",
},
{
name: "namespaced repo",
rkey: "org~myapp_v1.0.0",
wantRepository: "org/myapp",
wantTag: "v1.0.0",
},
{
name: "tag with underscore (splits on last underscore)",
rkey: "myapp_test_tag",
wantRepository: "myapp_test",
wantTag: "tag",
},
{
name: "deep namespace",
rkey: "a~b~c~myapp_prod",
wantRepository: "a/b/c/myapp",
wantTag: "prod",
},
{
name: "no underscore - all tag",
rkey: "latest",
wantRepository: "",
wantTag: "latest",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
gotRepo, gotTag := RKeyToRepositoryTag(tt.rkey)
if gotRepo != tt.wantRepository {
t.Errorf("RKeyToRepositoryTag() repository = %v, want %v", gotRepo, tt.wantRepository)
}
if gotTag != tt.wantTag {
t.Errorf("RKeyToRepositoryTag() tag = %v, want %v", gotTag, tt.wantTag)
}
})
}
}
// TestRepositoryTagRoundTrip tests that converting to rkey and back preserves values
// Note: Tags with underscores cannot be perfectly round-tripped since we use underscore as separator
func TestRepositoryTagRoundTrip(t *testing.T) {
tests := []struct {
repository string
tag string
}{
{"myapp", "latest"},
{"org/myapp", "v1.0.0"},
{"a/b/c/myapp", "prod"},
// Note: Tags with underscores are excluded - they cannot round-trip correctly
// because underscore is used as the separator between repository and tag
}
for _, tt := range tests {
t.Run(tt.repository+":"+tt.tag, func(t *testing.T) {
rkey := RepositoryTagToRKey(tt.repository, tt.tag)
gotRepo, gotTag := RKeyToRepositoryTag(rkey)
if gotRepo != tt.repository {
t.Errorf("Round trip failed: repository = %v, want %v", gotRepo, tt.repository)
}
if gotTag != tt.tag {
t.Errorf("Round trip failed: tag = %v, want %v", gotTag, tt.tag)
}
})
}
}
// TestNewManifestStore tests creating a new manifest store
func TestNewManifestStore(t *testing.T) {
client := NewClient("https://pds.example.com", "did:plc:test123", "token")
client := atproto.NewClient("https://pds.example.com", "did:plc:test123", "token")
blobStore := newMockBlobStore()
db := &mockDatabaseMetrics{}
store := NewManifestStore(
client,
"myapp",
"https://hold.example.com",
"did:web:hold.example.com",
"did:plc:alice123",
blobStore,
db,
)
ctx := mockRegistryContext(client, "myapp", "did:web:hold.example.com", "did:plc:alice123", "alice.test", db)
store := NewManifestStore(ctx, nil, blobStore)
if store.repository != "myapp" {
t.Errorf("repository = %v, want myapp", store.repository)
if store.ctx.Repository != "myapp" {
t.Errorf("repository = %v, want myapp", store.ctx.Repository)
}
if store.holdEndpoint != "https://hold.example.com" {
t.Errorf("holdEndpoint = %v, want https://hold.example.com", store.holdEndpoint)
if store.ctx.HoldDID != "did:web:hold.example.com" {
t.Errorf("holdDID = %v, want did:web:hold.example.com", store.ctx.HoldDID)
}
if store.holdDID != "did:web:hold.example.com" {
t.Errorf("holdDID = %v, want did:web:hold.example.com", store.holdDID)
if store.ctx.DID != "did:plc:alice123" {
t.Errorf("did = %v, want did:plc:alice123", store.ctx.DID)
}
if store.did != "did:plc:alice123" {
t.Errorf("did = %v, want did:plc:alice123", store.did)
if store.ctx.Handle != "alice.test" {
t.Errorf("handle = %v, want alice.test", store.ctx.Handle)
}
}
@@ -320,11 +187,12 @@ func TestManifestStore_GetLastFetchedHoldDID(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
client := NewClient("https://pds.example.com", "did:plc:test123", "token")
store := NewManifestStore(client, "myapp", "", "", "did:plc:test123", nil, nil)
client := atproto.NewClient("https://pds.example.com", "did:plc:test123", "token")
ctx := mockRegistryContext(client, "myapp", "", "did:plc:test123", "test.handle", nil)
store := NewManifestStore(ctx, nil, nil)
// Simulate what happens in Get() when parsing a manifest record
var manifestRecord ManifestRecord
var manifestRecord atproto.ManifestRecord
manifestRecord.HoldDID = tt.manifestHoldDID
manifestRecord.HoldEndpoint = tt.manifestHoldURL
@@ -332,7 +200,7 @@ func TestManifestStore_GetLastFetchedHoldDID(t *testing.T) {
if manifestRecord.HoldDID != "" {
store.lastFetchedHoldDID = manifestRecord.HoldDID
} else if manifestRecord.HoldEndpoint != "" {
store.lastFetchedHoldDID = ResolveHoldDIDFromURL(manifestRecord.HoldEndpoint)
store.lastFetchedHoldDID = atproto.ResolveHoldDIDFromURL(manifestRecord.HoldEndpoint)
}
got := store.GetLastFetchedHoldDID()
@@ -377,8 +245,8 @@ func TestRawManifest(t *testing.T) {
// TestExtractConfigLabels tests extracting labels from image config
func TestExtractConfigLabels(t *testing.T) {
// Create a mock config blob
configJSON := map[string]interface{}{
"config": map[string]interface{}{
configJSON := map[string]any{
"config": map[string]any{
"Labels": map[string]string{
"org.opencontainers.image.version": "1.0.0",
"org.opencontainers.image.authors": "test@example.com",
@@ -394,8 +262,9 @@ func TestExtractConfigLabels(t *testing.T) {
blobStore.blobs[configDigest] = configData
// Create manifest store
client := NewClient("https://pds.example.com", "did:plc:test123", "token")
store := NewManifestStore(client, "myapp", "", "", "did:plc:test123", blobStore, nil)
client := atproto.NewClient("https://pds.example.com", "did:plc:test123", "token")
ctx := mockRegistryContext(client, "myapp", "", "did:plc:test123", "test.handle", nil)
store := NewManifestStore(ctx, nil, blobStore)
// Extract labels
labels, err := store.extractConfigLabels(context.Background(), configDigest.String())
@@ -424,8 +293,8 @@ func TestExtractConfigLabels(t *testing.T) {
// TestExtractConfigLabels_NoLabels tests handling config without labels
func TestExtractConfigLabels_NoLabels(t *testing.T) {
// Config without Labels field
configJSON := map[string]interface{}{
"config": map[string]interface{}{},
configJSON := map[string]any{
"config": map[string]any{},
}
configData, _ := json.Marshal(configJSON)
@@ -433,8 +302,9 @@ func TestExtractConfigLabels_NoLabels(t *testing.T) {
configDigest := digest.FromBytes(configData)
blobStore.blobs[configDigest] = configData
client := NewClient("https://pds.example.com", "did:plc:test123", "token")
store := NewManifestStore(client, "myapp", "", "", "did:plc:test123", blobStore, nil)
client := atproto.NewClient("https://pds.example.com", "did:plc:test123", "token")
ctx := mockRegistryContext(client, "myapp", "", "did:plc:test123", "test.handle", nil)
store := NewManifestStore(ctx, nil, blobStore)
labels, err := store.extractConfigLabels(context.Background(), configDigest.String())
if err != nil {
@@ -450,8 +320,9 @@ func TestExtractConfigLabels_NoLabels(t *testing.T) {
// TestExtractConfigLabels_InvalidDigest tests error handling for invalid digest
func TestExtractConfigLabels_InvalidDigest(t *testing.T) {
blobStore := newMockBlobStore()
client := NewClient("https://pds.example.com", "did:plc:test123", "token")
store := NewManifestStore(client, "myapp", "", "", "did:plc:test123", blobStore, nil)
client := atproto.NewClient("https://pds.example.com", "did:plc:test123", "token")
ctx := mockRegistryContext(client, "myapp", "", "did:plc:test123", "test.handle", nil)
store := NewManifestStore(ctx, nil, blobStore)
_, err := store.extractConfigLabels(context.Background(), "invalid-digest")
if err == nil {
@@ -468,8 +339,9 @@ func TestExtractConfigLabels_InvalidJSON(t *testing.T) {
configDigest := digest.FromBytes(configData)
blobStore.blobs[configDigest] = configData
client := NewClient("https://pds.example.com", "did:plc:test123", "token")
store := NewManifestStore(client, "myapp", "", "", "did:plc:test123", blobStore, nil)
client := atproto.NewClient("https://pds.example.com", "did:plc:test123", "token")
ctx := mockRegistryContext(client, "myapp", "", "did:plc:test123", "test.handle", nil)
store := NewManifestStore(ctx, nil, blobStore)
_, err := store.extractConfigLabels(context.Background(), configDigest.String())
if err == nil {
@@ -480,18 +352,11 @@ func TestExtractConfigLabels_InvalidJSON(t *testing.T) {
// TestManifestStore_WithMetrics tests that metrics are tracked
func TestManifestStore_WithMetrics(t *testing.T) {
db := &mockDatabaseMetrics{}
client := NewClient("https://pds.example.com", "did:plc:test123", "token")
store := NewManifestStore(
client,
"myapp",
"https://hold.example.com",
"did:web:hold.example.com",
"did:plc:alice123",
nil,
db,
)
client := atproto.NewClient("https://pds.example.com", "did:plc:test123", "token")
ctx := mockRegistryContext(client, "myapp", "did:web:hold.example.com", "did:plc:alice123", "alice.test", db)
store := NewManifestStore(ctx, nil, nil)
if store.database != db {
if store.ctx.Database != db {
t.Error("ManifestStore should store database reference")
}
@@ -501,18 +366,11 @@ func TestManifestStore_WithMetrics(t *testing.T) {
// TestManifestStore_WithoutMetrics tests that nil database is acceptable
func TestManifestStore_WithoutMetrics(t *testing.T) {
client := NewClient("https://pds.example.com", "did:plc:test123", "token")
store := NewManifestStore(
client,
"myapp",
"https://hold.example.com",
"did:web:hold.example.com",
"did:plc:alice123",
nil,
nil, // nil database
)
client := atproto.NewClient("https://pds.example.com", "did:plc:test123", "token")
ctx := mockRegistryContext(client, "myapp", "did:web:hold.example.com", "did:plc:alice123", "alice.test", nil)
store := NewManifestStore(ctx, nil, nil)
if store.database != nil {
if store.ctx.Database != nil {
t.Error("ManifestStore should accept nil database")
}
}
+10 -11
View File
@@ -564,17 +564,16 @@ type CompletedPart struct {
// ProxyBlobWriter implements distribution.BlobWriter for proxy uploads using multipart upload
type ProxyBlobWriter struct {
store *ProxyBlobStore
options distribution.CreateOptions
uploadID string // S3 multipart upload ID
parts []CompletedPart // Track uploaded parts with ETags
partNumber int // Current part number (starts at 1)
buffer *bytes.Buffer // Buffer for current part
size int64 // Total bytes written
closed bool
id string // Distribution's upload ID (for state)
startedAt time.Time
finalDigest string // Set on Commit
store *ProxyBlobStore
options distribution.CreateOptions
uploadID string // S3 multipart upload ID
parts []CompletedPart // Track uploaded parts with ETags
partNumber int // Current part number (starts at 1)
buffer *bytes.Buffer // Buffer for current part
size int64 // Total bytes written
closed bool
id string // Distribution's upload ID (for state)
startedAt time.Time
}
// ID returns the upload ID
+1 -2
View File
@@ -313,8 +313,7 @@ func BenchmarkServiceTokenCacheAccess(b *testing.B) {
testTokenStr := "eyJhbGciOiJIUzI1NiJ9." + base64URLEncode(testPayload) + ".signature"
token.SetServiceToken(userDID, holdDID, testTokenStr)
b.ResetTimer()
for i := 0; i < b.N; i++ {
for b.Loop() {
cachedToken, expiresAt := token.GetServiceToken(userDID, holdDID)
if cachedToken == "" || time.Now().After(expiresAt) {
+62 -16
View File
@@ -2,10 +2,13 @@ package storage
import (
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"time"
"atcr.io/pkg/atproto"
"atcr.io/pkg/auth/oauth"
"github.com/distribution/distribution/v3"
)
@@ -13,9 +16,53 @@ import (
// The registry (AppView) is stateless and NEVER stores blobs locally
type RoutingRepository struct {
distribution.Repository
Ctx *RegistryContext // All context and services (exported for token updates)
manifestStore *atproto.ManifestStore // Cached manifest store instance
blobStore *ProxyBlobStore // Cached blob store instance
Ctx *RegistryContext // All context and services (exported for token updates)
manifestStore *ManifestStore // Cached manifest store instance
blobStore *ProxyBlobStore // Cached blob store instance
}
// refresherAdapter adapts the oauth.Refresher to implement atproto.HoldNotifier
type refresherAdapter struct {
refresher *oauth.Refresher
pdsEndpoint string
}
// GetServiceToken implements atproto.HoldNotifier
func (r *refresherAdapter) GetServiceToken(ctx context.Context, userDID, audienceDID string) (string, error) {
// Get OAuth session for the user
session, err := r.refresher.GetSession(ctx, userDID)
if err != nil {
return "", fmt.Errorf("failed to get OAuth session: %w", err)
}
// Build service auth URL
serviceAuthURL := fmt.Sprintf("%s/xrpc/com.atproto.server.getServiceAuth?aud=%s", r.pdsEndpoint, audienceDID)
req, err := http.NewRequestWithContext(ctx, "GET", serviceAuthURL, nil)
if err != nil {
return "", fmt.Errorf("failed to create request: %w", err)
}
// Use session's DoWithAuth to handle OAuth authentication automatically
resp, err := session.DoWithAuth(session.Client, req, "com.atproto.server.getServiceAuth")
if err != nil {
return "", fmt.Errorf("failed to request service token: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
body, _ := io.ReadAll(resp.Body)
return "", fmt.Errorf("PDS returned status %d: %s", resp.StatusCode, body)
}
var result struct {
Token string `json:"token"`
}
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
return "", fmt.Errorf("failed to decode response: %w", err)
}
return result.Token, nil
}
// NewRoutingRepository creates a new routing repository
@@ -33,17 +80,16 @@ func (r *RoutingRepository) Manifests(ctx context.Context, options ...distributi
// Ensure blob store is created first (needed for label extraction during push)
blobStore := r.Blobs(ctx)
// ManifestStore needs both DID and URL for backward compat (legacy holdEndpoint field)
// For now, pass holdDID twice (will be cleaned up in manifest_store.go later)
r.manifestStore = atproto.NewManifestStore(
r.Ctx.ATProtoClient,
r.Ctx.Repository,
r.Ctx.HoldDID,
r.Ctx.HoldDID,
r.Ctx.DID,
blobStore,
r.Ctx.Database,
)
// Wrap the Refresher in an adapter to implement HoldNotifier
var notifier HoldNotifier
if r.Ctx.Refresher != nil {
notifier = &refresherAdapter{
refresher: r.Ctx.Refresher,
pdsEndpoint: r.Ctx.PDSEndpoint,
}
}
r.manifestStore = NewManifestStore(r.Ctx, notifier, blobStore)
}
// After any manifest operation, cache the hold DID for blob fetches
@@ -102,5 +148,5 @@ func (r *RoutingRepository) Blobs(ctx context.Context) distribution.BlobStore {
// Tags returns the tag service
// Tags are stored in ATProto as io.atcr.tag records
func (r *RoutingRepository) Tags(ctx context.Context) distribution.TagService {
return atproto.NewTagStore(r.Ctx.ATProtoClient, r.Ctx.Repository)
return NewTagStore(r.Ctx.ATProtoClient, r.Ctx.Repository)
}
@@ -1,10 +1,11 @@
package atproto
package storage
import (
"context"
"encoding/json"
"fmt"
"atcr.io/pkg/atproto"
"github.com/distribution/distribution/v3"
"github.com/opencontainers/go-digest"
)
@@ -12,12 +13,12 @@ import (
// TagStore implements distribution.TagService
// It stores tags in ATProto as records
type TagStore struct {
client *Client
client *atproto.Client
repository string
}
// NewTagStore creates a new ATProto-backed tag store
func NewTagStore(client *Client, repository string) *TagStore {
func NewTagStore(client *atproto.Client, repository string) *TagStore {
return &TagStore{
client: client,
repository: repository,
@@ -27,15 +28,15 @@ func NewTagStore(client *Client, repository string) *TagStore {
// Get retrieves the descriptor for a tag
func (s *TagStore) Get(ctx context.Context, tag string) (distribution.Descriptor, error) {
// Build record key
rkey := RepositoryTagToRKey(s.repository, tag)
rkey := atproto.RepositoryTagToRKey(s.repository, tag)
// Fetch tag record from ATProto
record, err := s.client.GetRecord(ctx, TagCollection, rkey)
record, err := s.client.GetRecord(ctx, atproto.TagCollection, rkey)
if err != nil {
return distribution.Descriptor{}, distribution.ErrTagUnknown{Tag: tag}
}
var tagRecord TagRecord
var tagRecord atproto.TagRecord
if err := json.Unmarshal(record.Value, &tagRecord); err != nil {
return distribution.Descriptor{}, fmt.Errorf("failed to unmarshal tag record: %w", err)
}
@@ -62,11 +63,11 @@ func (s *TagStore) Get(ctx context.Context, tag string) (distribution.Descriptor
// Tag associates a tag with a descriptor (manifest digest)
func (s *TagStore) Tag(ctx context.Context, tag string, desc distribution.Descriptor) error {
// Create tag record with manifest AT-URI
tagRecord := NewTagRecord(s.client.DID(), s.repository, tag, desc.Digest.String())
tagRecord := atproto.NewTagRecord(s.client.DID(), s.repository, tag, desc.Digest.String())
// Store in ATProto
rkey := RepositoryTagToRKey(s.repository, tag)
_, err := s.client.PutRecord(ctx, TagCollection, rkey, tagRecord)
rkey := atproto.RepositoryTagToRKey(s.repository, tag)
_, err := s.client.PutRecord(ctx, atproto.TagCollection, rkey, tagRecord)
if err != nil {
return fmt.Errorf("failed to store tag in ATProto: %w", err)
}
@@ -76,21 +77,21 @@ func (s *TagStore) Tag(ctx context.Context, tag string, desc distribution.Descri
// Untag removes a tag
func (s *TagStore) Untag(ctx context.Context, tag string) error {
rkey := RepositoryTagToRKey(s.repository, tag)
return s.client.DeleteRecord(ctx, TagCollection, rkey)
rkey := atproto.RepositoryTagToRKey(s.repository, tag)
return s.client.DeleteRecord(ctx, atproto.TagCollection, rkey)
}
// All returns all tags for this repository
func (s *TagStore) All(ctx context.Context) ([]string, error) {
// List all records in the tag collection
records, err := s.client.ListRecords(ctx, TagCollection, 100)
records, err := s.client.ListRecords(ctx, atproto.TagCollection, 100)
if err != nil {
return nil, fmt.Errorf("failed to list tags: %w", err)
}
var tags []string
for _, record := range records {
var tagRecord TagRecord
var tagRecord atproto.TagRecord
if err := json.Unmarshal(record.Value, &tagRecord); err != nil {
// Skip invalid records
continue
@@ -108,14 +109,14 @@ func (s *TagStore) All(ctx context.Context) ([]string, error) {
// Lookup returns the set of tags for a given digest
func (s *TagStore) Lookup(ctx context.Context, desc distribution.Descriptor) ([]string, error) {
// List all records in the tag collection
records, err := s.client.ListRecords(ctx, TagCollection, 100)
records, err := s.client.ListRecords(ctx, atproto.TagCollection, 100)
if err != nil {
return nil, fmt.Errorf("failed to list tags: %w", err)
}
var tags []string
for _, record := range records {
var tagRecord TagRecord
var tagRecord atproto.TagRecord
if err := json.Unmarshal(record.Value, &tagRecord); err != nil {
// Skip invalid records
continue
@@ -1,4 +1,4 @@
package atproto
package storage
import (
"context"
@@ -8,13 +8,14 @@ import (
"strings"
"testing"
"atcr.io/pkg/atproto"
"github.com/distribution/distribution/v3"
"github.com/opencontainers/go-digest"
)
// TestNewTagStore tests creating a new tag store
func TestNewTagStore(t *testing.T) {
client := NewClient("https://pds.example.com", "did:plc:test123", "token")
client := atproto.NewClient("https://pds.example.com", "did:plc:test123", "token")
store := NewTagStore(client, "myapp")
if store.repository != "myapp" {
@@ -67,12 +68,12 @@ func TestTagStore_Get(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// Verify query parameters
query := r.URL.Query()
rkey := RepositoryTagToRKey("myapp", tt.tag)
rkey := atproto.RepositoryTagToRKey("myapp", tt.tag)
if query.Get("rkey") != rkey {
t.Errorf("rkey = %v, want %v", query.Get("rkey"), rkey)
}
if query.Get("collection") != TagCollection {
t.Errorf("collection = %v, want %v", query.Get("collection"), TagCollection)
if query.Get("collection") != atproto.TagCollection {
t.Errorf("collection = %v, want %v", query.Get("collection"), atproto.TagCollection)
}
w.WriteHeader(tt.serverStatus)
@@ -80,7 +81,7 @@ func TestTagStore_Get(t *testing.T) {
}))
defer server.Close()
client := NewClient(server.URL, "did:plc:test123", "test-token")
client := atproto.NewClient(server.URL, "did:plc:test123", "test-token")
store := NewTagStore(client, "myapp")
desc, err := store.Get(context.Background(), tt.tag)
@@ -119,7 +120,7 @@ func TestTagStore_Get_InvalidDigest(t *testing.T) {
}))
defer server.Close()
client := NewClient(server.URL, "did:plc:test123", "test-token")
client := atproto.NewClient(server.URL, "did:plc:test123", "test-token")
store := NewTagStore(client, "myapp")
_, err := store.Get(context.Background(), "latest")
@@ -148,7 +149,7 @@ func TestTagStore_Get_BackwardCompatibility(t *testing.T) {
}))
defer server.Close()
client := NewClient(server.URL, "did:plc:test123", "test-token")
client := atproto.NewClient(server.URL, "did:plc:test123", "test-token")
store := NewTagStore(client, "myapp")
desc, err := store.Get(context.Background(), "latest")
@@ -181,7 +182,7 @@ func TestTagStore_Get_NewManifestField(t *testing.T) {
}))
defer server.Close()
client := NewClient(server.URL, "did:plc:test123", "test-token")
client := atproto.NewClient(server.URL, "did:plc:test123", "test-token")
store := NewTagStore(client, "myapp")
desc, err := store.Get(context.Background(), "latest")
@@ -228,7 +229,7 @@ func TestTagStore_Tag(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var sentTagRecord *TagRecord
var sentTagRecord *atproto.TagRecord
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != "POST" {
@@ -236,24 +237,24 @@ func TestTagStore_Tag(t *testing.T) {
}
// Parse request body
var body map[string]interface{}
var body map[string]any
json.NewDecoder(r.Body).Decode(&body)
// Verify rkey
expectedRKey := RepositoryTagToRKey("myapp", tt.tag)
expectedRKey := atproto.RepositoryTagToRKey("myapp", tt.tag)
if body["rkey"] != expectedRKey {
t.Errorf("rkey = %v, want %v", body["rkey"], expectedRKey)
}
// Verify collection
if body["collection"] != TagCollection {
t.Errorf("collection = %v, want %v", body["collection"], TagCollection)
if body["collection"] != atproto.TagCollection {
t.Errorf("collection = %v, want %v", body["collection"], atproto.TagCollection)
}
// Parse and verify tag record
recordData := body["record"].(map[string]interface{})
recordData := body["record"].(map[string]any)
recordBytes, _ := json.Marshal(recordData)
var tagRecord TagRecord
var tagRecord atproto.TagRecord
json.Unmarshal(recordBytes, &tagRecord)
sentTagRecord = &tagRecord
@@ -266,7 +267,7 @@ func TestTagStore_Tag(t *testing.T) {
}))
defer server.Close()
client := NewClient(server.URL, "did:plc:test123", "test-token")
client := atproto.NewClient(server.URL, "did:plc:test123", "test-token")
store := NewTagStore(client, "myapp")
desc := distribution.Descriptor{
@@ -283,8 +284,8 @@ func TestTagStore_Tag(t *testing.T) {
if !tt.wantErr && sentTagRecord != nil {
// Verify the tag record
if sentTagRecord.Type != TagCollection {
t.Errorf("Type = %v, want %v", sentTagRecord.Type, TagCollection)
if sentTagRecord.Type != atproto.TagCollection {
t.Errorf("Type = %v, want %v", sentTagRecord.Type, atproto.TagCollection)
}
if sentTagRecord.Repository != "myapp" {
t.Errorf("Repository = %v, want myapp", sentTagRecord.Repository)
@@ -293,7 +294,7 @@ func TestTagStore_Tag(t *testing.T) {
t.Errorf("Tag = %v, want %v", sentTagRecord.Tag, tt.tag)
}
// New records should have manifest field
expectedURI := BuildManifestURI("did:plc:test123", tt.digest.String())
expectedURI := atproto.BuildManifestURI("did:plc:test123", tt.digest.String())
if sentTagRecord.Manifest != expectedURI {
t.Errorf("Manifest = %v, want %v", sentTagRecord.Manifest, expectedURI)
}
@@ -337,10 +338,10 @@ func TestTagStore_Untag(t *testing.T) {
}
// Parse body to verify delete parameters
var body map[string]interface{}
var body map[string]any
json.NewDecoder(r.Body).Decode(&body)
expectedRKey := RepositoryTagToRKey("myapp", tt.tag)
expectedRKey := atproto.RepositoryTagToRKey("myapp", tt.tag)
if body["rkey"] != expectedRKey {
t.Errorf("rkey = %v, want %v", body["rkey"], expectedRKey)
}
@@ -354,7 +355,7 @@ func TestTagStore_Untag(t *testing.T) {
}))
defer server.Close()
client := NewClient(server.URL, "did:plc:test123", "test-token")
client := atproto.NewClient(server.URL, "did:plc:test123", "test-token")
store := NewTagStore(client, "myapp")
err := store.Untag(context.Background(), tt.tag)
@@ -422,8 +423,8 @@ func TestTagStore_All(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// Verify query parameters
query := r.URL.Query()
if query.Get("collection") != TagCollection {
t.Errorf("collection = %v, want %v", query.Get("collection"), TagCollection)
if query.Get("collection") != atproto.TagCollection {
t.Errorf("collection = %v, want %v", query.Get("collection"), atproto.TagCollection)
}
if query.Get("limit") != "100" {
t.Errorf("limit = %v, want 100", query.Get("limit"))
@@ -434,7 +435,7 @@ func TestTagStore_All(t *testing.T) {
}))
defer server.Close()
client := NewClient(server.URL, "did:plc:test123", "test-token")
client := atproto.NewClient(server.URL, "did:plc:test123", "test-token")
store := NewTagStore(client, "myapp")
tags, err := store.All(context.Background())
@@ -496,7 +497,7 @@ func TestTagStore_All_SkipsInvalidRecords(t *testing.T) {
}))
defer server.Close()
client := NewClient(server.URL, "did:plc:test123", "test-token")
client := atproto.NewClient(server.URL, "did:plc:test123", "test-token")
store := NewTagStore(client, "myapp")
tags, err := store.All(context.Background())
@@ -584,7 +585,7 @@ func TestTagStore_Lookup(t *testing.T) {
}))
defer server.Close()
client := NewClient(server.URL, "did:plc:test123", "test-token")
client := atproto.NewClient(server.URL, "did:plc:test123", "test-token")
store := NewTagStore(client, "myapp")
desc := distribution.Descriptor{
@@ -646,7 +647,7 @@ func TestTagStore_Lookup_FiltersByRepository(t *testing.T) {
}))
defer server.Close()
client := NewClient(server.URL, "did:plc:test123", "test-token")
client := atproto.NewClient(server.URL, "did:plc:test123", "test-token")
store := NewTagStore(client, "myapp") // Looking for "myapp" tags only
desc := distribution.Descriptor{
@@ -676,7 +677,7 @@ func TestTagStore_ListRecordsError(t *testing.T) {
}))
defer server.Close()
client := NewClient(server.URL, "did:plc:test123", "test-token")
client := atproto.NewClient(server.URL, "did:plc:test123", "test-token")
store := NewTagStore(client, "myapp")
// Test All()
@@ -700,7 +701,7 @@ func TestTagStore_GetErrorTypes(t *testing.T) {
}))
defer server.Close()
client := NewClient(server.URL, "did:plc:test123", "test-token")
client := atproto.NewClient(server.URL, "did:plc:test123", "test-token")
store := NewTagStore(client, "myapp")
_, err := store.Get(context.Background(), "notfound")
+1 -1
View File
@@ -607,7 +607,7 @@ func TestTemplateExecution_WithFuncMap(t *testing.T) {
tests := []struct {
name string
templateStr string
data interface{}
data any
expectInOutput string
}{
{
+388 -2
View File
@@ -300,7 +300,7 @@ func (t *CaptainRecord) MarshalCBOR(w io.Writer) error {
}
cw := cbg.NewCborWriter(w)
fieldCount := 7
fieldCount := 8
if t.Region == "" {
fieldCount--
@@ -466,6 +466,22 @@ func (t *CaptainRecord) MarshalCBOR(w io.Writer) error {
if err := cbg.WriteBool(w, t.AllowAllCrew); err != nil {
return err
}
// t.EnableManifestPosts (bool) (bool)
if len("enableManifestPosts") > 8192 {
return xerrors.Errorf("Value in field \"enableManifestPosts\" was too long")
}
if err := cw.WriteMajorTypeHeader(cbg.MajTextString, uint64(len("enableManifestPosts"))); err != nil {
return err
}
if _, err := cw.WriteString(string("enableManifestPosts")); err != nil {
return err
}
if err := cbg.WriteBool(w, t.EnableManifestPosts); err != nil {
return err
}
return nil
}
@@ -494,7 +510,7 @@ func (t *CaptainRecord) UnmarshalCBOR(r io.Reader) (err error) {
n := extra
nameBuf := make([]byte, 12)
nameBuf := make([]byte, 19)
for i := uint64(0); i < n; i++ {
nameLen, ok, err := cbg.ReadFullStringIntoBuf(cr, nameBuf, 8192)
if err != nil {
@@ -601,6 +617,376 @@ func (t *CaptainRecord) UnmarshalCBOR(r io.Reader) (err error) {
default:
return fmt.Errorf("booleans are either major type 7, value 20 or 21 (got %d)", extra)
}
// t.EnableManifestPosts (bool) (bool)
case "enableManifestPosts":
maj, extra, err = cr.ReadHeader()
if err != nil {
return err
}
if maj != cbg.MajOther {
return fmt.Errorf("booleans must be major type 7")
}
switch extra {
case 20:
t.EnableManifestPosts = false
case 21:
t.EnableManifestPosts = true
default:
return fmt.Errorf("booleans are either major type 7, value 20 or 21 (got %d)", extra)
}
default:
// Field doesn't exist on this type, so ignore it
if err := cbg.ScanForLinks(r, func(cid.Cid) {}); err != nil {
return err
}
}
}
return nil
}
func (t *LayerRecord) MarshalCBOR(w io.Writer) error {
if t == nil {
_, err := w.Write(cbg.CborNull)
return err
}
cw := cbg.NewCborWriter(w)
if _, err := cw.Write([]byte{168}); err != nil {
return err
}
// t.Size (int64) (int64)
if len("size") > 8192 {
return xerrors.Errorf("Value in field \"size\" was too long")
}
if err := cw.WriteMajorTypeHeader(cbg.MajTextString, uint64(len("size"))); err != nil {
return err
}
if _, err := cw.WriteString(string("size")); err != nil {
return err
}
if t.Size >= 0 {
if err := cw.WriteMajorTypeHeader(cbg.MajUnsignedInt, uint64(t.Size)); err != nil {
return err
}
} else {
if err := cw.WriteMajorTypeHeader(cbg.MajNegativeInt, uint64(-t.Size-1)); err != nil {
return err
}
}
// t.Type (string) (string)
if len("$type") > 8192 {
return xerrors.Errorf("Value in field \"$type\" was too long")
}
if err := cw.WriteMajorTypeHeader(cbg.MajTextString, uint64(len("$type"))); err != nil {
return err
}
if _, err := cw.WriteString(string("$type")); err != nil {
return err
}
if len(t.Type) > 8192 {
return xerrors.Errorf("Value in field t.Type was too long")
}
if err := cw.WriteMajorTypeHeader(cbg.MajTextString, uint64(len(t.Type))); err != nil {
return err
}
if _, err := cw.WriteString(string(t.Type)); err != nil {
return err
}
// t.Digest (string) (string)
if len("digest") > 8192 {
return xerrors.Errorf("Value in field \"digest\" was too long")
}
if err := cw.WriteMajorTypeHeader(cbg.MajTextString, uint64(len("digest"))); err != nil {
return err
}
if _, err := cw.WriteString(string("digest")); err != nil {
return err
}
if len(t.Digest) > 8192 {
return xerrors.Errorf("Value in field t.Digest was too long")
}
if err := cw.WriteMajorTypeHeader(cbg.MajTextString, uint64(len(t.Digest))); err != nil {
return err
}
if _, err := cw.WriteString(string(t.Digest)); err != nil {
return err
}
// t.UserDID (string) (string)
if len("userDid") > 8192 {
return xerrors.Errorf("Value in field \"userDid\" was too long")
}
if err := cw.WriteMajorTypeHeader(cbg.MajTextString, uint64(len("userDid"))); err != nil {
return err
}
if _, err := cw.WriteString(string("userDid")); err != nil {
return err
}
if len(t.UserDID) > 8192 {
return xerrors.Errorf("Value in field t.UserDID was too long")
}
if err := cw.WriteMajorTypeHeader(cbg.MajTextString, uint64(len(t.UserDID))); err != nil {
return err
}
if _, err := cw.WriteString(string(t.UserDID)); err != nil {
return err
}
// t.CreatedAt (string) (string)
if len("createdAt") > 8192 {
return xerrors.Errorf("Value in field \"createdAt\" was too long")
}
if err := cw.WriteMajorTypeHeader(cbg.MajTextString, uint64(len("createdAt"))); err != nil {
return err
}
if _, err := cw.WriteString(string("createdAt")); err != nil {
return err
}
if len(t.CreatedAt) > 8192 {
return xerrors.Errorf("Value in field t.CreatedAt was too long")
}
if err := cw.WriteMajorTypeHeader(cbg.MajTextString, uint64(len(t.CreatedAt))); err != nil {
return err
}
if _, err := cw.WriteString(string(t.CreatedAt)); err != nil {
return err
}
// t.MediaType (string) (string)
if len("mediaType") > 8192 {
return xerrors.Errorf("Value in field \"mediaType\" was too long")
}
if err := cw.WriteMajorTypeHeader(cbg.MajTextString, uint64(len("mediaType"))); err != nil {
return err
}
if _, err := cw.WriteString(string("mediaType")); err != nil {
return err
}
if len(t.MediaType) > 8192 {
return xerrors.Errorf("Value in field t.MediaType was too long")
}
if err := cw.WriteMajorTypeHeader(cbg.MajTextString, uint64(len(t.MediaType))); err != nil {
return err
}
if _, err := cw.WriteString(string(t.MediaType)); err != nil {
return err
}
// t.Repository (string) (string)
if len("repository") > 8192 {
return xerrors.Errorf("Value in field \"repository\" was too long")
}
if err := cw.WriteMajorTypeHeader(cbg.MajTextString, uint64(len("repository"))); err != nil {
return err
}
if _, err := cw.WriteString(string("repository")); err != nil {
return err
}
if len(t.Repository) > 8192 {
return xerrors.Errorf("Value in field t.Repository was too long")
}
if err := cw.WriteMajorTypeHeader(cbg.MajTextString, uint64(len(t.Repository))); err != nil {
return err
}
if _, err := cw.WriteString(string(t.Repository)); err != nil {
return err
}
// t.UserHandle (string) (string)
if len("userHandle") > 8192 {
return xerrors.Errorf("Value in field \"userHandle\" was too long")
}
if err := cw.WriteMajorTypeHeader(cbg.MajTextString, uint64(len("userHandle"))); err != nil {
return err
}
if _, err := cw.WriteString(string("userHandle")); err != nil {
return err
}
if len(t.UserHandle) > 8192 {
return xerrors.Errorf("Value in field t.UserHandle was too long")
}
if err := cw.WriteMajorTypeHeader(cbg.MajTextString, uint64(len(t.UserHandle))); err != nil {
return err
}
if _, err := cw.WriteString(string(t.UserHandle)); err != nil {
return err
}
return nil
}
func (t *LayerRecord) UnmarshalCBOR(r io.Reader) (err error) {
*t = LayerRecord{}
cr := cbg.NewCborReader(r)
maj, extra, err := cr.ReadHeader()
if err != nil {
return err
}
defer func() {
if err == io.EOF {
err = io.ErrUnexpectedEOF
}
}()
if maj != cbg.MajMap {
return fmt.Errorf("cbor input should be of type map")
}
if extra > cbg.MaxLength {
return fmt.Errorf("LayerRecord: map struct too large (%d)", extra)
}
n := extra
nameBuf := make([]byte, 10)
for i := uint64(0); i < n; i++ {
nameLen, ok, err := cbg.ReadFullStringIntoBuf(cr, nameBuf, 8192)
if err != nil {
return err
}
if !ok {
// Field doesn't exist on this type, so ignore it
if err := cbg.ScanForLinks(cr, func(cid.Cid) {}); err != nil {
return err
}
continue
}
switch string(nameBuf[:nameLen]) {
// t.Size (int64) (int64)
case "size":
{
maj, extra, err := cr.ReadHeader()
if err != nil {
return err
}
var extraI int64
switch maj {
case cbg.MajUnsignedInt:
extraI = int64(extra)
if extraI < 0 {
return fmt.Errorf("int64 positive overflow")
}
case cbg.MajNegativeInt:
extraI = int64(extra)
if extraI < 0 {
return fmt.Errorf("int64 negative overflow")
}
extraI = -1 - extraI
default:
return fmt.Errorf("wrong type for int64 field: %d", maj)
}
t.Size = int64(extraI)
}
// t.Type (string) (string)
case "$type":
{
sval, err := cbg.ReadStringWithMax(cr, 8192)
if err != nil {
return err
}
t.Type = string(sval)
}
// t.Digest (string) (string)
case "digest":
{
sval, err := cbg.ReadStringWithMax(cr, 8192)
if err != nil {
return err
}
t.Digest = string(sval)
}
// t.UserDID (string) (string)
case "userDid":
{
sval, err := cbg.ReadStringWithMax(cr, 8192)
if err != nil {
return err
}
t.UserDID = string(sval)
}
// t.CreatedAt (string) (string)
case "createdAt":
{
sval, err := cbg.ReadStringWithMax(cr, 8192)
if err != nil {
return err
}
t.CreatedAt = string(sval)
}
// t.MediaType (string) (string)
case "mediaType":
{
sval, err := cbg.ReadStringWithMax(cr, 8192)
if err != nil {
return err
}
t.MediaType = string(sval)
}
// t.Repository (string) (string)
case "repository":
{
sval, err := cbg.ReadStringWithMax(cr, 8192)
if err != nil {
return err
}
t.Repository = string(sval)
}
// t.UserHandle (string) (string)
case "userHandle":
{
sval, err := cbg.ReadStringWithMax(cr, 8192)
if err != nil {
return err
}
t.UserHandle = string(sval)
}
default:
// Field doesn't exist on this type, so ignore it
+4 -4
View File
@@ -34,7 +34,7 @@ func TestPutRecord(t *testing.T) {
name string
collection string
rkey string
record interface{}
record any
serverResponse string
serverStatus int
wantErr bool
@@ -93,7 +93,7 @@ func TestPutRecord(t *testing.T) {
}
// Verify request body
var body map[string]interface{}
var body map[string]any
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
t.Errorf("Failed to decode request body: %v", err)
}
@@ -158,7 +158,7 @@ func TestGetRecord(t *testing.T) {
t.Errorf("URI = %v, want at://did:plc:test123/io.atcr.manifest/abc123", r.URI)
}
var value map[string]interface{}
var value map[string]any
if err := json.Unmarshal(r.Value, &value); err != nil {
t.Errorf("Failed to unmarshal value: %v", err)
}
@@ -290,7 +290,7 @@ func TestDeleteRecord(t *testing.T) {
}
// Verify request body
var body map[string]interface{}
var body map[string]any
if err := json.NewDecoder(r.Body).Decode(&body); err != nil {
t.Errorf("Failed to decode request body: %v", err)
}
+6
View File
@@ -39,6 +39,12 @@ const (
// Request: {"uploadId": "..."}
// Response: {"status": "aborted"}
HoldAbortUpload = "/xrpc/io.atcr.hold.abortUpload"
// HoldNotifyManifest notifies hold about a manifest upload for layer tracking and Bluesky posting.
// Method: POST
// Request: {"repository": "...", "tag": "...", "userDid": "...", "userHandle": "...", "manifest": {...}}
// Response: {"success": true, "layersCreated": 5, "postCreated": true, "postUri": "at://..."}
HoldNotifyManifest = "/xrpc/io.atcr.hold.notifyManifest"
)
// Hold service crew management endpoints (io.atcr.hold.*)
+2 -1
View File
@@ -25,10 +25,11 @@ import (
)
func main() {
// Generate map-style encoders for CrewRecord, CaptainRecord, and TangledProfileRecord
// Generate map-style encoders for CrewRecord, CaptainRecord, LayerRecord, and TangledProfileRecord
if err := cbg.WriteMapEncodersToFile("cbor_gen.go", "atproto",
atproto.CrewRecord{},
atproto.CaptainRecord{},
atproto.LayerRecord{},
atproto.TangledProfileRecord{},
); err != nil {
fmt.Printf("Failed to generate CBOR encoders: %v\n", err)
+75 -7
View File
@@ -34,6 +34,10 @@ const (
// Note: Uses same collection name as HoldCrewCollection but stored in different PDS (hold's PDS vs owner's PDS)
CrewCollection = "io.atcr.hold.crew"
// LayerCollection is the collection name for container layer metadata
// Stored in hold's embedded PDS to track which layers are stored
LayerCollection = "io.atcr.hold.layer"
// TangledProfileCollection is the collection name for tangled profiles
// Stored in hold's embedded PDS (singleton record at rkey "self")
TangledProfileCollection = "sh.tangled.actor.profile"
@@ -434,6 +438,40 @@ func isDID(s string) bool {
return len(s) > 4 && s[:4] == "did:"
}
// RepositoryTagToRKey converts a repository and tag to an ATProto record key
// ATProto record keys must match: ^[a-zA-Z0-9._~-]{1,512}$
func RepositoryTagToRKey(repository, tag string) string {
// Combine repository and tag to create a unique key
// Replace invalid characters: slashes become tildes (~)
// We use tilde instead of dash to avoid ambiguity with repository names that contain hyphens
key := fmt.Sprintf("%s_%s", repository, tag)
// Replace / with ~ (slash not allowed in rkeys, tilde is allowed and unlikely in repo names)
key = strings.ReplaceAll(key, "/", "~")
return key
}
// RKeyToRepositoryTag converts an ATProto record key back to repository and tag
// This is the inverse of RepositoryTagToRKey
// Note: If the tag contains underscores, this will split on the LAST underscore
func RKeyToRepositoryTag(rkey string) (repository, tag string) {
// Find the last underscore to split repository and tag
lastUnderscore := strings.LastIndex(rkey, "_")
if lastUnderscore == -1 {
// No underscore found - treat entire string as tag with empty repository
return "", rkey
}
repository = rkey[:lastUnderscore]
tag = rkey[lastUnderscore+1:]
// Convert tildes back to slashes in repository (tilde was used to encode slashes)
repository = strings.ReplaceAll(repository, "~", "/")
return repository, tag
}
// BuildManifestURI creates an AT-URI for a manifest record
// did: The DID of the user (e.g., "did:plc:xyz123")
// manifestDigest: The manifest digest (e.g., "sha256:abc123...")
@@ -498,13 +536,14 @@ func (t *TagRecord) GetManifestDigest() (string, error) {
// Stored in the hold's embedded PDS to identify the hold owner and settings
// Uses CBOR encoding for efficient storage in hold's carstore
type CaptainRecord struct {
Type string `json:"$type" cborgen:"$type"`
Owner string `json:"owner" cborgen:"owner"` // DID of hold owner
Public bool `json:"public" cborgen:"public"` // Public read access
AllowAllCrew bool `json:"allowAllCrew" cborgen:"allowAllCrew"` // Allow any authenticated user to register as crew
DeployedAt string `json:"deployedAt" cborgen:"deployedAt"` // RFC3339 timestamp
Region string `json:"region,omitempty" cborgen:"region,omitempty"` // S3 region (optional)
Provider string `json:"provider,omitempty" cborgen:"provider,omitempty"` // Deployment provider (optional)
Type string `json:"$type" cborgen:"$type"`
Owner string `json:"owner" cborgen:"owner"` // DID of hold owner
Public bool `json:"public" cborgen:"public"` // Public read access
AllowAllCrew bool `json:"allowAllCrew" cborgen:"allowAllCrew"` // Allow any authenticated user to register as crew
EnableManifestPosts bool `json:"enableManifestPosts" cborgen:"enableManifestPosts"` // Enable Bluesky posts when manifests are pushed (overrides env var)
DeployedAt string `json:"deployedAt" cborgen:"deployedAt"` // RFC3339 timestamp
Region string `json:"region,omitempty" cborgen:"region,omitempty"` // S3 region (optional)
Provider string `json:"provider,omitempty" cborgen:"provider,omitempty"` // Deployment provider (optional)
}
// CrewRecord represents a crew member in the hold
@@ -520,6 +559,35 @@ type CrewRecord struct {
AddedAt string `json:"addedAt" cborgen:"addedAt"` // RFC3339 timestamp
}
// LayerRecord represents metadata about a container layer stored in the hold
// Collection: io.atcr.hold.layer
// Stored in the hold's embedded PDS for tracking and analytics
// Uses CBOR encoding for efficient storage in hold's carstore
type LayerRecord struct {
Type string `json:"$type" cborgen:"$type"`
Digest string `json:"digest" cborgen:"digest"` // Layer digest (e.g., "sha256:abc123...")
Size int64 `json:"size" cborgen:"size"` // Size in bytes
MediaType string `json:"mediaType" cborgen:"mediaType"` // Media type (e.g., "application/vnd.oci.image.layer.v1.tar+gzip")
Repository string `json:"repository" cborgen:"repository"` // Repository this layer belongs to
UserDID string `json:"userDid" cborgen:"userDid"` // DID of user who uploaded this layer
UserHandle string `json:"userHandle" cborgen:"userHandle"` // Handle of user (for display purposes)
CreatedAt string `json:"createdAt" cborgen:"createdAt"` // RFC3339 timestamp
}
// NewLayerRecord creates a new layer record
func NewLayerRecord(digest string, size int64, mediaType, repository, userDID, userHandle string) *LayerRecord {
return &LayerRecord{
Type: LayerCollection,
Digest: digest,
Size: size,
MediaType: mediaType,
Repository: repository,
UserDID: userDID,
UserHandle: userHandle,
CreatedAt: time.Now().Format(time.RFC3339),
}
}
// TangledProfileRecord represents a Tangled profile for the hold
// Collection: sh.tangled.actor.profile (singleton record at rkey "self")
// Stored in the hold's embedded PDS
+7 -7
View File
@@ -48,11 +48,11 @@ func TestEnsureProfile_Create(t *testing.T) {
// Second request: PutRecord (create profile)
if r.Method == "POST" && strings.Contains(r.URL.Path, "putRecord") {
var body map[string]interface{}
var body map[string]any
json.NewDecoder(r.Body).Decode(&body)
// Verify profile data
recordData := body["record"].(map[string]interface{})
recordData := body["record"].(map[string]any)
if recordData["$type"] != SailorProfileCollection {
t.Errorf("$type = %v, want %v", recordData["$type"], SailorProfileCollection)
}
@@ -218,7 +218,7 @@ func TestGetProfile(t *testing.T) {
migrationLocks = sync.Map{}
putRecordCalled := false
var migrationRequest map[string]interface{}
var migrationRequest map[string]any
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// GetRecord
@@ -273,7 +273,7 @@ func TestGetProfile(t *testing.T) {
}
if migrationRequest != nil {
recordData := migrationRequest["record"].(map[string]interface{})
recordData := migrationRequest["record"].(map[string]any)
migratedHold := recordData["defaultHold"]
if migratedHold != tt.expectedHoldDID {
t.Errorf("Migrated defaultHold = %v, want %v", migratedHold, tt.expectedHoldDID)
@@ -401,11 +401,11 @@ func TestUpdateProfile(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var sentProfile map[string]interface{}
var sentProfile map[string]any
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method == "POST" && strings.Contains(r.URL.Path, "putRecord") {
var body map[string]interface{}
var body map[string]any
json.NewDecoder(r.Body).Decode(&body)
sentProfile = body
@@ -432,7 +432,7 @@ func TestUpdateProfile(t *testing.T) {
if !tt.wantErr {
// Verify normalization happened
recordData := sentProfile["record"].(map[string]interface{})
recordData := sentProfile["record"].(map[string]any)
defaultHold := recordData["defaultHold"]
// Handle empty string (may be nil in JSON)
defaultHoldStr := ""
+2 -3
View File
@@ -4,6 +4,7 @@ import (
"context"
"encoding/json"
"fmt"
"maps"
"os"
"path/filepath"
"sync"
@@ -178,9 +179,7 @@ func (s *FileStore) ListSessions() map[string]*oauth.ClientSessionData {
// Return a copy to prevent external modification
result := make(map[string]*oauth.ClientSessionData)
for k, v := range s.sessions {
result[k] = v
}
maps.Copy(result, s.sessions)
return result
}
-1
View File
@@ -107,7 +107,6 @@ func (v *SessionValidator) CreateSessionAndGetToken(ctx context.Context, identif
return "", "", "", fmt.Errorf("failed to resolve identity %q: %w", identifier, err)
}
did = ident.DID.String()
pds := ident.PDSEndpoint()
if pds == "" {
return "", "", "", fmt.Errorf("no PDS endpoint found for %q", identifier)
+2 -2
View File
@@ -123,7 +123,7 @@ func InvalidateServiceToken(did, holdDID string) {
}
// GetCacheStats returns statistics about the service token cache for debugging
func GetCacheStats() map[string]interface{} {
func GetCacheStats() map[string]any {
globalServiceTokensMu.RLock()
defer globalServiceTokensMu.RUnlock()
@@ -139,7 +139,7 @@ func GetCacheStats() map[string]interface{} {
}
}
return map[string]interface{}{
return map[string]any{
"total_entries": len(globalServiceTokens),
"valid_tokens": validCount,
"expired_tokens": expiredCount,
+118
View File
@@ -46,6 +46,7 @@ func (h *XRPCHandler) RegisterHandlers(r chi.Router) {
r.Put(atproto.HoldUploadPart, h.HandleUploadPart)
r.Post(atproto.HoldCompleteUpload, h.HandleCompleteUpload)
r.Post(atproto.HoldAbortUpload, h.HandleAbortUpload)
r.Post(atproto.HoldNotifyManifest, h.HandleNotifyManifest)
})
}
@@ -197,6 +198,123 @@ func (h *XRPCHandler) HandleAbortUpload(w http.ResponseWriter, r *http.Request)
})
}
// HandleNotifyManifest handles manifest upload notifications from AppView
// Creates layer records and optionally posts to Bluesky
func (h *XRPCHandler) HandleNotifyManifest(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
// Validate service token (same auth as blob:write endpoints)
validatedUser, err := pds.ValidateBlobWriteAccess(r, h.pds, h.httpClient)
if err != nil {
RespondError(w, http.StatusForbidden, fmt.Sprintf("authorization failed: %v", err))
return
}
// Parse request
var req struct {
Repository string `json:"repository"`
Tag string `json:"tag"`
UserDID string `json:"userDid"`
UserHandle string `json:"userHandle"`
Manifest struct {
MediaType string `json:"mediaType"`
Config struct {
Digest string `json:"digest"`
Size int64 `json:"size"`
} `json:"config"`
Layers []struct {
Digest string `json:"digest"`
Size int64 `json:"size"`
MediaType string `json:"mediaType"`
} `json:"layers"`
} `json:"manifest"`
}
if err := DecodeJSON(r, &req); err != nil {
RespondError(w, http.StatusBadRequest, err.Error())
return
}
// Verify user DID matches token
if req.UserDID != validatedUser.DID {
RespondError(w, http.StatusForbidden, "user DID mismatch")
return
}
// Check if manifest posts are enabled
// TODO: Check captain record enableManifestPosts field
// For now, posts are always created
postsEnabled := true
// Create layer records for each blob
layersCreated := 0
for _, layer := range req.Manifest.Layers {
record := atproto.NewLayerRecord(
layer.Digest,
layer.Size,
layer.MediaType,
req.Repository,
req.UserDID,
req.UserHandle,
)
_, _, err := h.pds.CreateLayerRecord(ctx, record)
if err != nil {
fmt.Printf("Failed to create layer record: %v\n", err)
// Continue creating other records
} else {
layersCreated++
}
}
// Calculate total size from all layers
var totalSize int64
for _, layer := range req.Manifest.Layers {
totalSize += layer.Size
}
totalSize += req.Manifest.Config.Size // Add config blob size
// Create Bluesky post if enabled
var postURI string
postCreated := false
if postsEnabled {
// Extract manifest digest from first layer (or use config digest as fallback)
manifestDigest := req.Manifest.Config.Digest
if len(req.Manifest.Layers) > 0 {
manifestDigest = req.Manifest.Layers[0].Digest
}
postURI, err = h.pds.CreateManifestPost(
ctx,
req.Repository,
req.Tag,
req.UserHandle,
manifestDigest,
totalSize,
)
if err != nil {
fmt.Printf("Failed to create manifest post: %v\n", err)
} else {
postCreated = true
}
}
// Return response
resp := map[string]any{
"success": layersCreated > 0 || postCreated,
"layersCreated": layersCreated,
"postCreated": postCreated,
}
if postURI != "" {
resp["postUri"] = postURI
}
if err != nil && layersCreated == 0 && !postCreated {
resp["error"] = err.Error()
}
RespondJSON(w, http.StatusOK, resp)
}
// requireBlobWriteAccess middleware - validates DPoP + OAuth and checks for blob:write permission
func (h *XRPCHandler) requireBlobWriteAccess(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+2 -6
View File
@@ -480,7 +480,7 @@ func ValidateServiceToken(r *http.Request, holdDID string, httpClient HTTPClient
}
// Fetch public key from issuer's DID document
publicKey, err := fetchPublicKeyFromDID(r.Context(), issuerDID, httpClient)
publicKey, err := fetchPublicKeyFromDID(r.Context(), issuerDID)
if err != nil {
return nil, fmt.Errorf("failed to fetch public key for issuer %s: %w", issuerDID, err)
}
@@ -502,11 +502,7 @@ func ValidateServiceToken(r *http.Request, holdDID string, httpClient HTTPClient
// fetchPublicKeyFromDID fetches the public key from a DID document
// Supports did:plc and did:web
// Returns the atcrypto.PublicKey for signature verification
func fetchPublicKeyFromDID(ctx context.Context, did string, httpClient HTTPClient) (atcrypto.PublicKey, error) {
if httpClient == nil {
httpClient = http.DefaultClient
}
func fetchPublicKeyFromDID(ctx context.Context, did string) (atcrypto.PublicKey, error) {
// Use indigo's identity resolution
directory := identity.DefaultDirectory()
atID, err := syntax.ParseAtIdentifier(did)
-28
View File
@@ -239,34 +239,6 @@ func (h *ServiceTokenTestHelper) AddServiceTokenToRequest(req *http.Request, exp
return nil
}
// mockDIDResolver is a simple mock for DID resolution that returns a fixed public key
type mockDIDResolver struct {
publicKeys map[string]atcrypto.PublicKey
}
// newMockDIDResolver creates a new mock DID resolver
func newMockDIDResolver() *mockDIDResolver {
return &mockDIDResolver{
publicKeys: make(map[string]atcrypto.PublicKey),
}
}
// RegisterDID registers a DID with its public key
func (m *mockDIDResolver) RegisterDID(did string, publicKey atcrypto.PublicKey) {
m.publicKeys[did] = publicKey
}
// Do implements the HTTPClient interface for mocking DID resolution
// This intercepts fetchPublicKeyFromDID's indigo directory calls
func (m *mockDIDResolver) Do(req *http.Request) (*http.Response, error) {
// This mock is not used directly - we'll need to inject the public key differently
// For now, return a 404 to indicate DID resolution should use our registered keys
return &http.Response{
StatusCode: http.StatusNotFound,
Body: http.NoBody,
}, nil
}
// TestValidateServiceToken_ValidToken tests validation of a properly formed service token
func TestValidateServiceToken_ValidToken(t *testing.T) {
// This test validates token structure, audience, and expiration
+3 -3
View File
@@ -29,9 +29,9 @@ type EventBroadcaster struct {
eventSeq int64
eventHistory []HistoricalEvent // Ring buffer for cursor backfill (deprecated, kept for compatibility)
maxHistory int
holdDID string // DID of the hold for setting repo field
db *sql.DB // Database for persistent event storage
dbPath string // Path to database file
holdDID string // DID of the hold for setting repo field
db *sql.DB // Database for persistent event storage
dbPath string // Path to database file
}
// Subscriber represents a WebSocket client subscribed to the firehose
+59
View File
@@ -0,0 +1,59 @@
package pds
import (
"context"
"fmt"
"atcr.io/pkg/atproto"
)
// CreateLayerRecord creates a new layer record in the hold's PDS
// Returns the rkey and CID of the created record
func (p *HoldPDS) CreateLayerRecord(ctx context.Context, record *atproto.LayerRecord) (string, string, error) {
// Validate record
if record.Type != atproto.LayerCollection {
return "", "", fmt.Errorf("invalid record type: %s", record.Type)
}
if record.Digest == "" {
return "", "", fmt.Errorf("digest is required")
}
if record.Size <= 0 {
return "", "", fmt.Errorf("size must be positive")
}
// Create record with auto-generated TID rkey
rkey, recordCID, err := p.repomgr.CreateRecord(
ctx,
p.uid,
atproto.LayerCollection,
record,
)
if err != nil {
return "", "", fmt.Errorf("failed to create layer record: %w", err)
}
return rkey, recordCID.String(), nil
}
// GetLayerRecord retrieves a specific layer record by rkey
// Note: This is a simplified implementation. For production, you may need to pass the CID
func (p *HoldPDS) GetLayerRecord(ctx context.Context, rkey string) (*atproto.LayerRecord, error) {
// For now, we don't implement this as it's not needed for the manifest post feature
// Full implementation would require querying the carstore with a specific CID
return nil, fmt.Errorf("GetLayerRecord not yet implemented - use via XRPC listRecords instead")
}
// ListLayerRecords lists layer records with pagination
// Returns records, next cursor (empty if no more), and error
// Note: This is a simplified implementation. For production, consider adding filters
// (by repository, user, digest, etc.) and proper pagination
func (p *HoldPDS) ListLayerRecords(ctx context.Context, limit int, cursor string) ([]*atproto.LayerRecord, string, error) {
// For now, return empty list - full implementation would query the carstore
// This would require iterating over records in the collection and filtering
// In practice, layer records are mainly for analytics and Bluesky posts,
// not for runtime queries
return nil, "", fmt.Errorf("ListLayerRecords not yet implemented")
}
+294
View File
@@ -0,0 +1,294 @@
package pds
import (
"testing"
"atcr.io/pkg/atproto"
)
func TestCreateLayerRecord(t *testing.T) {
// Setup test PDS
pds, ctx := setupTestPDS(t)
tests := []struct {
name string
record *atproto.LayerRecord
wantErr bool
errSubstr string
}{
{
name: "valid layer record",
record: atproto.NewLayerRecord(
"sha256:abc123def456",
1048576, // 1 MB
"application/vnd.oci.image.layer.v1.tar+gzip",
"myapp",
"did:plc:alice123",
"alice.bsky.social",
),
wantErr: false,
},
{
name: "valid layer record with large size",
record: atproto.NewLayerRecord(
"sha256:fedcba987654",
1073741824, // 1 GB
"application/vnd.docker.image.rootfs.diff.tar.gzip",
"debian",
"did:plc:bob456",
"bob.example.com",
),
wantErr: false,
},
{
name: "invalid record type",
record: &atproto.LayerRecord{
Type: "wrong.type",
Digest: "sha256:abc123",
Size: 1024,
MediaType: "application/vnd.oci.image.layer.v1.tar",
Repository: "test",
UserDID: "did:plc:test",
UserHandle: "test.example.com",
},
wantErr: true,
errSubstr: "invalid record type",
},
{
name: "missing digest",
record: &atproto.LayerRecord{
Type: atproto.LayerCollection,
Digest: "",
Size: 1024,
MediaType: "application/vnd.oci.image.layer.v1.tar",
Repository: "test",
UserDID: "did:plc:test",
UserHandle: "test.example.com",
},
wantErr: true,
errSubstr: "digest is required",
},
{
name: "zero size",
record: &atproto.LayerRecord{
Type: atproto.LayerCollection,
Digest: "sha256:abc123",
Size: 0,
MediaType: "application/vnd.oci.image.layer.v1.tar",
Repository: "test",
UserDID: "did:plc:test",
UserHandle: "test.example.com",
},
wantErr: true,
errSubstr: "size must be positive",
},
{
name: "negative size",
record: &atproto.LayerRecord{
Type: atproto.LayerCollection,
Digest: "sha256:abc123",
Size: -1,
MediaType: "application/vnd.oci.image.layer.v1.tar",
Repository: "test",
UserDID: "did:plc:test",
UserHandle: "test.example.com",
},
wantErr: true,
errSubstr: "size must be positive",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
rkey, cid, err := pds.CreateLayerRecord(ctx, tt.record)
if tt.wantErr {
if err == nil {
t.Errorf("CreateLayerRecord() expected error containing %q, got nil", tt.errSubstr)
return
}
if tt.errSubstr != "" && !contains(err.Error(), tt.errSubstr) {
t.Errorf("CreateLayerRecord() error = %v, want error containing %q", err, tt.errSubstr)
}
return
}
if err != nil {
t.Errorf("CreateLayerRecord() unexpected error: %v", err)
return
}
if rkey == "" {
t.Error("CreateLayerRecord() returned empty rkey")
}
if cid == "" {
t.Error("CreateLayerRecord() returned empty CID")
}
t.Logf("Created layer record: rkey=%s, cid=%s", rkey, cid)
})
}
}
func TestCreateLayerRecord_MultipleRecords(t *testing.T) {
// Test creating multiple layer records for the same manifest
pds, ctx := setupTestPDS(t)
layers := []struct {
digest string
size int64
}{
{"sha256:layer1abc123", 1024},
{"sha256:layer2def456", 2048},
{"sha256:layer3ghi789", 4096},
}
createdRKeys := make(map[string]bool)
for i, layer := range layers {
record := atproto.NewLayerRecord(
layer.digest,
layer.size,
"application/vnd.oci.image.layer.v1.tar+gzip",
"multi-layer-app",
"did:plc:test123",
"test.example.com",
)
rkey, cid, err := pds.CreateLayerRecord(ctx, record)
if err != nil {
t.Fatalf("CreateLayerRecord() for layer %d failed: %v", i, err)
}
// Ensure unique rkeys
if createdRKeys[rkey] {
t.Errorf("CreateLayerRecord() returned duplicate rkey: %s", rkey)
}
createdRKeys[rkey] = true
t.Logf("Layer %d: rkey=%s, cid=%s", i, rkey, cid)
}
if len(createdRKeys) != len(layers) {
t.Errorf("Created %d unique rkeys, want %d", len(createdRKeys), len(layers))
}
}
func TestNewLayerRecord(t *testing.T) {
// Test the layer record constructor
digest := "sha256:abc123def456"
size := int64(1048576)
mediaType := "application/vnd.oci.image.layer.v1.tar+gzip"
repository := "myapp"
userDID := "did:plc:alice123"
userHandle := "alice.bsky.social"
record := atproto.NewLayerRecord(digest, size, mediaType, repository, userDID, userHandle)
if record == nil {
t.Fatal("NewLayerRecord() returned nil")
}
// Verify all fields are set correctly
if record.Type != atproto.LayerCollection {
t.Errorf("Type = %q, want %q", record.Type, atproto.LayerCollection)
}
if record.Digest != digest {
t.Errorf("Digest = %q, want %q", record.Digest, digest)
}
if record.Size != size {
t.Errorf("Size = %d, want %d", record.Size, size)
}
if record.MediaType != mediaType {
t.Errorf("MediaType = %q, want %q", record.MediaType, mediaType)
}
if record.Repository != repository {
t.Errorf("Repository = %q, want %q", record.Repository, repository)
}
if record.UserDID != userDID {
t.Errorf("UserDID = %q, want %q", record.UserDID, userDID)
}
if record.UserHandle != userHandle {
t.Errorf("UserHandle = %q, want %q", record.UserHandle, userHandle)
}
if record.CreatedAt == "" {
t.Error("CreatedAt is empty")
}
t.Logf("Created layer record: %+v", record)
}
func TestLayerRecord_FieldValidation(t *testing.T) {
// Test various field values
tests := []struct {
name string
digest string
size int64
mediaType string
repository string
userDID string
userHandle string
}{
{
name: "typical OCI layer",
digest: "sha256:e692418e4cbaf90ca69d05a66403747baa33ee08806650b51fab815ad7fc331f",
size: 12582912, // 12 MB
mediaType: "application/vnd.oci.image.layer.v1.tar+gzip",
repository: "hsm-secrets-operator",
userDID: "did:plc:evan123",
userHandle: "evan.jarrett.net",
},
{
name: "Docker layer format",
digest: "sha256:abc123",
size: 1024,
mediaType: "application/vnd.docker.image.rootfs.diff.tar.gzip",
repository: "nginx",
userDID: "did:plc:user456",
userHandle: "user.example.com",
},
{
name: "uncompressed layer",
digest: "sha256:def456",
size: 2048,
mediaType: "application/vnd.oci.image.layer.v1.tar",
repository: "alpine",
userDID: "did:plc:user789",
userHandle: "user.bsky.social",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
record := atproto.NewLayerRecord(
tt.digest,
tt.size,
tt.mediaType,
tt.repository,
tt.userDID,
tt.userHandle,
)
if record == nil {
t.Fatal("NewLayerRecord() returned nil")
}
// Verify the record can be created
if record.Type != atproto.LayerCollection {
t.Errorf("Type = %q, want %q", record.Type, atproto.LayerCollection)
}
if record.Digest != tt.digest {
t.Errorf("Digest = %q, want %q", record.Digest, tt.digest)
}
})
}
}
+150
View File
@@ -0,0 +1,150 @@
package pds
import (
"context"
"fmt"
"strings"
"time"
bsky "github.com/bluesky-social/indigo/api/bsky"
)
// CreateManifestPost creates a Bluesky post announcing a manifest upload
// Includes facets for clickable mentions and links
func (p *HoldPDS) CreateManifestPost(
ctx context.Context,
repository, tag, userHandle, digest string,
totalSize int64,
) (string, error) {
now := time.Now()
// Build AppView repository URL
appViewURL := fmt.Sprintf("https://atcr.io/r/%s/%s", userHandle, repository)
// Format post text components
digestShort := formatDigest(digest)
sizeStr := formatSize(totalSize)
repoWithTag := fmt.Sprintf("%s:%s", repository, tag)
// Build text: "@alice.bsky.social just pushed hsm-secrets-operator:latest\nDigest: sha256:abc...def Size: 12.2 MB"
text := fmt.Sprintf("@%s just pushed %s\nDigest: %s Size: %s", userHandle, repoWithTag, digestShort, sizeStr)
// Create facets for mentions and links
facets := buildFacets(text, userHandle, repoWithTag, appViewURL)
// Create post struct with facets
post := &bsky.FeedPost{
LexiconTypeID: "app.bsky.feed.post",
Text: text,
Facets: facets,
CreatedAt: now.Format(time.RFC3339),
}
// Create record with auto-generated TID
rkey, recordCID, err := p.repomgr.CreateRecord(
ctx,
p.uid,
"app.bsky.feed.post",
post,
)
if err != nil {
return "", fmt.Errorf("failed to create manifest post: %w", err)
}
// Build ATProto URI for the post
postURI := fmt.Sprintf("at://%s/app.bsky.feed.post/%s", p.did, rkey)
fmt.Printf("Created manifest post: %s (cid: %s)\n", postURI, recordCID)
return postURI, nil
}
// formatDigest truncates digest to first 7 and last 7 chars
// Example: sha256:abc1234567890...fedcba9876543210 -> sha256:abc1234...9876543
func formatDigest(digest string) string {
if !strings.HasPrefix(digest, "sha256:") {
return digest // Return as-is if not sha256
}
hash := strings.TrimPrefix(digest, "sha256:")
if len(hash) <= 14 {
return digest // Too short to truncate
}
return fmt.Sprintf("sha256:%s...%s", hash[:7], hash[len(hash)-7:])
}
// formatSize converts bytes to human-readable format
// Examples: 1024 -> "1.0 KB", 1048576 -> "1.0 MB", 1073741824 -> "1.0 GB"
func formatSize(bytes int64) string {
const (
KB = 1024
MB = 1024 * KB
GB = 1024 * MB
)
switch {
case bytes >= GB:
return fmt.Sprintf("%.1f GB", float64(bytes)/float64(GB))
case bytes >= MB:
return fmt.Sprintf("%.1f MB", float64(bytes)/float64(MB))
case bytes >= KB:
return fmt.Sprintf("%.1f KB", float64(bytes)/float64(KB))
default:
return fmt.Sprintf("%d B", bytes)
}
}
// buildFacets creates mention and link facets for rich text
// IMPORTANT: Byte offsets must be calculated for UTF-8 encoded text
func buildFacets(text, userHandle, repoWithTag, appViewURL string) []*bsky.RichtextFacet {
facets := []*bsky.RichtextFacet{}
// Find mention: "@alice.bsky.social"
mentionText := "@" + userHandle
mentionStart := strings.Index(text, mentionText)
if mentionStart >= 0 {
// Calculate byte offsets (not character offsets!)
byteStart := int64(len(text[:mentionStart]))
byteEnd := int64(len(text[:mentionStart+len(mentionText)]))
facets = append(facets, &bsky.RichtextFacet{
Index: &bsky.RichtextFacet_ByteSlice{
ByteStart: byteStart,
ByteEnd: byteEnd,
},
Features: []*bsky.RichtextFacet_Features_Elem{
{
RichtextFacet_Mention: &bsky.RichtextFacet_Mention{
Did: "", // Will be resolved by Bluesky from handle
},
},
},
})
}
// Find repository link: "hsm-secrets-operator:latest"
linkStart := strings.Index(text, repoWithTag)
if linkStart >= 0 {
// Calculate byte offsets
byteStart := int64(len(text[:linkStart]))
byteEnd := int64(len(text[:linkStart+len(repoWithTag)]))
facets = append(facets, &bsky.RichtextFacet{
Index: &bsky.RichtextFacet_ByteSlice{
ByteStart: byteStart,
ByteEnd: byteEnd,
},
Features: []*bsky.RichtextFacet_Features_Elem{
{
RichtextFacet_Link: &bsky.RichtextFacet_Link{
Uri: appViewURL,
},
},
},
})
}
return facets
}
+335
View File
@@ -0,0 +1,335 @@
package pds
import (
"strings"
"testing"
bsky "github.com/bluesky-social/indigo/api/bsky"
)
func TestFormatDigest(t *testing.T) {
tests := []struct {
name string
digest string
expected string
}{
{
name: "standard sha256 digest",
digest: "sha256:abc1234567890fedcba9876543210",
expected: "sha256:abc1234...6543210", // Last 7 chars of hash
},
{
name: "short digest (no truncation)",
digest: "sha256:abc123",
expected: "sha256:abc123",
},
{
name: "non-sha256 digest",
digest: "sha512:abc123",
expected: "sha512:abc123",
},
{
name: "real sha256 digest",
digest: "sha256:e692418e4cbaf90ca69d05a66403747baa33ee08806650b51fab815ad7fc331f",
expected: "sha256:e692418...7fc331f", // Last 7 chars are "7fc331f"
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := formatDigest(tt.digest)
if result != tt.expected {
t.Errorf("formatDigest(%q) = %q, want %q", tt.digest, result, tt.expected)
}
})
}
}
func TestFormatSize(t *testing.T) {
tests := []struct {
name string
bytes int64
expected string
}{
{
name: "bytes",
bytes: 512,
expected: "512 B",
},
{
name: "kilobytes",
bytes: 1024,
expected: "1.0 KB",
},
{
name: "kilobytes with decimal",
bytes: 1536, // 1.5 KB
expected: "1.5 KB",
},
{
name: "megabytes",
bytes: 1048576, // 1 MB
expected: "1.0 MB",
},
{
name: "megabytes with decimal",
bytes: 12582912, // 12 MB
expected: "12.0 MB",
},
{
name: "gigabytes",
bytes: 1073741824, // 1 GB
expected: "1.0 GB",
},
{
name: "gigabytes with decimal",
bytes: 2147483648, // 2 GB
expected: "2.0 GB",
},
{
name: "zero bytes",
bytes: 0,
expected: "0 B",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := formatSize(tt.bytes)
if result != tt.expected {
t.Errorf("formatSize(%d) = %q, want %q", tt.bytes, result, tt.expected)
}
})
}
}
func TestBuildFacets(t *testing.T) {
tests := []struct {
name string
text string
userHandle string
repoWithTag string
appViewURL string
wantFacets int // number of facets expected
}{
{
name: "standard post with mention and link",
text: "@alice.bsky.social just pushed myapp:latest\nDigest: sha256:abc...def Size: 12.2 MB",
userHandle: "alice.bsky.social",
repoWithTag: "myapp:latest",
appViewURL: "https://atcr.io/r/alice.bsky.social/myapp",
wantFacets: 2,
},
{
name: "no matches found",
text: "random text",
userHandle: "alice.bsky.social",
repoWithTag: "myapp:latest",
appViewURL: "https://atcr.io/r/alice.bsky.social/myapp",
wantFacets: 0,
},
{
name: "only mention found",
text: "@alice.bsky.social did something",
userHandle: "alice.bsky.social",
repoWithTag: "myapp:latest",
appViewURL: "https://atcr.io/r/alice.bsky.social/myapp",
wantFacets: 1,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
facets := buildFacets(tt.text, tt.userHandle, tt.repoWithTag, tt.appViewURL)
if len(facets) != tt.wantFacets {
t.Errorf("buildFacets() returned %d facets, want %d", len(facets), tt.wantFacets)
}
// Verify facet structure for standard case
if tt.name == "standard post with mention and link" && len(facets) == 2 {
// Check mention facet
mentionFacet := facets[0]
if mentionFacet.Index == nil {
t.Error("mention facet has nil Index")
}
if len(mentionFacet.Features) != 1 {
t.Errorf("mention facet has %d features, want 1", len(mentionFacet.Features))
}
if mentionFacet.Features[0].RichtextFacet_Mention == nil {
t.Error("mention facet feature is not a mention")
}
// Check link facet
linkFacet := facets[1]
if linkFacet.Index == nil {
t.Error("link facet has nil Index")
}
if len(linkFacet.Features) != 1 {
t.Errorf("link facet has %d features, want 1", len(linkFacet.Features))
}
if linkFacet.Features[0].RichtextFacet_Link == nil {
t.Error("link facet feature is not a link")
}
if linkFacet.Features[0].RichtextFacet_Link.Uri != tt.appViewURL {
t.Errorf("link facet URI = %q, want %q", linkFacet.Features[0].RichtextFacet_Link.Uri, tt.appViewURL)
}
}
})
}
}
func TestBuildFacets_ByteOffsets(t *testing.T) {
// Test that byte offsets are correctly calculated
text := "@alice.bsky.social just pushed myapp:latest"
userHandle := "alice.bsky.social"
repoWithTag := "myapp:latest"
appViewURL := "https://atcr.io/r/alice.bsky.social/myapp"
facets := buildFacets(text, userHandle, repoWithTag, appViewURL)
if len(facets) != 2 {
t.Fatalf("expected 2 facets, got %d", len(facets))
}
// Check mention facet byte offsets
mentionFacet := facets[0]
mentionText := "@alice.bsky.social"
expectedStart := int64(0) // mention is at the start
expectedEnd := int64(len(mentionText))
if mentionFacet.Index.ByteStart != expectedStart {
t.Errorf("mention ByteStart = %d, want %d", mentionFacet.Index.ByteStart, expectedStart)
}
if mentionFacet.Index.ByteEnd != expectedEnd {
t.Errorf("mention ByteEnd = %d, want %d", mentionFacet.Index.ByteEnd, expectedEnd)
}
// Verify the mention text extraction
extractedMention := text[mentionFacet.Index.ByteStart:mentionFacet.Index.ByteEnd]
if extractedMention != mentionText {
t.Errorf("extracted mention = %q, want %q", extractedMention, mentionText)
}
// Check link facet byte offsets
linkFacet := facets[1]
linkStart := len("@alice.bsky.social just pushed ")
expectedLinkStart := int64(linkStart)
expectedLinkEnd := int64(linkStart + len(repoWithTag))
if linkFacet.Index.ByteStart != expectedLinkStart {
t.Errorf("link ByteStart = %d, want %d", linkFacet.Index.ByteStart, expectedLinkStart)
}
if linkFacet.Index.ByteEnd != expectedLinkEnd {
t.Errorf("link ByteEnd = %d, want %d", linkFacet.Index.ByteEnd, expectedLinkEnd)
}
// Verify the link text extraction
extractedLink := text[linkFacet.Index.ByteStart:linkFacet.Index.ByteEnd]
if extractedLink != repoWithTag {
t.Errorf("extracted link = %q, want %q", extractedLink, repoWithTag)
}
}
func TestBuildFacets_UTF8Handling(t *testing.T) {
// Test with Unicode characters to ensure byte offsets work correctly
text := "@alice.bsky.social just pushed 🚀myapp:latest"
userHandle := "alice.bsky.social"
repoWithTag := "🚀myapp:latest" // Note: emoji is multi-byte
appViewURL := "https://atcr.io/r/alice.bsky.social/myapp"
facets := buildFacets(text, userHandle, repoWithTag, appViewURL)
if len(facets) != 2 {
t.Fatalf("expected 2 facets, got %d", len(facets))
}
// Verify that byte extraction works with UTF-8
mentionFacet := facets[0]
extractedMention := text[mentionFacet.Index.ByteStart:mentionFacet.Index.ByteEnd]
expectedMention := "@alice.bsky.social"
if extractedMention != expectedMention {
t.Errorf("extracted mention = %q, want %q", extractedMention, expectedMention)
}
linkFacet := facets[1]
extractedLink := text[linkFacet.Index.ByteStart:linkFacet.Index.ByteEnd]
if extractedLink != repoWithTag {
t.Errorf("extracted link = %q, want %q", extractedLink, repoWithTag)
}
}
func TestBuildFacets_NoOverlap(t *testing.T) {
// Ensure facets don't overlap
text := "@alice.bsky.social just pushed myapp:latest"
userHandle := "alice.bsky.social"
repoWithTag := "myapp:latest"
appViewURL := "https://atcr.io/r/alice.bsky.social/myapp"
facets := buildFacets(text, userHandle, repoWithTag, appViewURL)
if len(facets) != 2 {
t.Fatalf("expected 2 facets, got %d", len(facets))
}
// Facets should not overlap
facet1 := facets[0]
facet2 := facets[1]
if facet1.Index.ByteEnd > facet2.Index.ByteStart {
t.Errorf("facets overlap: facet1 ends at %d, facet2 starts at %d",
facet1.Index.ByteEnd, facet2.Index.ByteStart)
}
}
func TestBuildFacets_RealWorldExample(t *testing.T) {
// Test with the actual example from the requirements
repository := "hsm-secrets-operator"
tag := "latest"
userHandle := "evan.jarrett.net"
digest := "sha256:e692418e4cbaf90ca69d05a66403747baa33ee08806650b51fab815ad7fc331f"
totalSize := int64(12800000) // ~12.2 MB
repoWithTag := repository + ":" + tag
digestShort := formatDigest(digest)
sizeStr := formatSize(totalSize)
text := "@" + userHandle + " just pushed " + repoWithTag + "\nDigest: " + digestShort + " Size: " + sizeStr
appViewURL := "https://atcr.io/r/" + userHandle + "/" + repository
facets := buildFacets(text, userHandle, repoWithTag, appViewURL)
// Should have 2 facets: mention and link
if len(facets) != 2 {
t.Fatalf("expected 2 facets, got %d", len(facets))
}
// Verify the complete post structure
post := &bsky.FeedPost{
LexiconTypeID: "app.bsky.feed.post",
Text: text,
Facets: facets,
}
if post.Text == "" {
t.Error("post text is empty")
}
if len(post.Facets) != 2 {
t.Errorf("post has %d facets, want 2", len(post.Facets))
}
// Verify text contains expected components
expectedTexts := []string{
"@" + userHandle,
repoWithTag,
digestShort,
sizeStr,
}
for _, expected := range expectedTexts {
if !strings.Contains(text, expected) {
t.Errorf("post text missing expected component: %q", expected)
}
}
}
+2 -2
View File
@@ -139,7 +139,7 @@ type userLock struct {
}
func (rm *RepoManager) lockUser(ctx context.Context, user models.Uid) func() {
ctx, span := otel.Tracer("repoman").Start(ctx, "userLock")
_, span := otel.Tracer("repoman").Start(ctx, "userLock")
defer span.End()
rm.lklk.Lock()
@@ -1062,7 +1062,7 @@ func (rm *RepoManager) ImportNewRepo(ctx context.Context, user models.Uid, repoD
return nil
})
if err != nil {
return fmt.Errorf("process new repo (current rev: %s): %w:", currev, err)
return fmt.Errorf("process new repo (current rev: %s): %w", currev, err)
}
return nil
+2 -1
View File
@@ -20,10 +20,11 @@ import (
// init registers our custom ATProto types with indigo's lexutil type registry
// This allows repomgr.GetRecord to automatically unmarshal our types
func init() {
// Register captain, crew, and tangled profile record types
// Register captain, crew, tangled profile, and layer record types
// These must match the $type field in the records
lexutil.RegisterType(atproto.CaptainCollection, &atproto.CaptainRecord{})
lexutil.RegisterType(atproto.CrewCollection, &atproto.CrewRecord{})
lexutil.RegisterType(atproto.LayerCollection, &atproto.LayerRecord{})
lexutil.RegisterType(atproto.TangledProfileCollection, &atproto.TangledProfileRecord{})
}
+1 -1
View File
@@ -339,7 +339,7 @@ func (h *XRPCHandler) HandleGetProfiles(w http.ResponseWriter, r *http.Request)
// buildProfileResponse builds a profile response map (shared by GetProfile and GetProfiles)
func (h *XRPCHandler) buildProfileResponse(ctx context.Context) map[string]any {
// Get profile record from repo
_, profileVal, err := h.pds.repomgr.GetRecord(
_, profileVal, _ := h.pds.repomgr.GetRecord(
ctx,
h.pds.uid,
"app.bsky.actor.profile",