mirror of
https://tangled.org/evan.jarrett.net/at-container-registry
synced 2026-09-20 17:24:16 +00:00
post to bluesky when manifests uploaded. linting fixes
This commit is contained in:
@@ -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)]
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(¬ifyResp); 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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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")
|
||||
@@ -607,7 +607,7 @@ func TestTemplateExecution_WithFuncMap(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
templateStr string
|
||||
data interface{}
|
||||
data any
|
||||
expectInOutput string
|
||||
}{
|
||||
{
|
||||
|
||||
+388
-2
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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.*)
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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 := ""
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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{})
|
||||
}
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user