diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..70eb643 --- /dev/null +++ b/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2026 wisp-deploy contributors + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/README.md b/README.md index 20b0402..e934ce9 100644 --- a/README.md +++ b/README.md @@ -1,28 +1,64 @@ # wisp-deploy -minimal go version of [wispcli](https://tangled.org/did:plc:f5yyql4gsffx7jjplezhttle/tree/main/cli) for CI. +`wisp-deploy` is a small, non-interactive go cli for deploying a directory to [wisp.place](https://wisp.place) from ci. it resolves an at protocol handle, finds its pds, uploads changed files, and writes the `place.wisp.*` records used by wisp.place. -for interactive usage, i recommend using wispcli as it provides a nicer visual experience and importantly, oauth login. +it only deploys sites. use [wispctl](https://tangled.org/nekomimi.pet/wisp.place-monorepo/tree/main/cli) if you want oauth login or the other interactive commands. -```bash -Usage: wisp-deploy [options] +## authentication -Deploy a static site to wisp.place +create an app password for the account and put it in the environment under `WISP_APP_PASSWORD`. this is the preferred option because the password does not appear in the process arguments. -Options: - -p, --path Directory to deploy - -s, --site Site name (defaults to directory name) - --directory Enable directory listing - --spa Enable SPA mode (serve index.html for all routes) - -c, --concurrency Number of concurrent uploads (backs off to 2 on rate limit) (default: 3) - --force-gzip Force gzip compression for all files regardless of type - --password App password for headless authentication - -h, --help display help for command +`--password` is kept for compatibility with wispctl. command arguments may be visible to other processes or recorded by a ci runner, so avoid it when the environment variable is available. `--password-file` is also supported. an explicit flag or password file takes precedence over `WISP_APP_PASSWORD`; `--password` and `--password-file` cannot be used together. + +handles are normally resolved through their `/.well-known/atproto-did` endpoint, with the `_atproto` dns record as a fallback. the resulting `did:plc` or `did:web` document provides the pds endpoint. pds connections must use https. + +set `WISP_MINIDOC_URL` to the full `blue.microcosm.identity.resolveMiniDoc` xrpc endpoint to use a slingshot instance instead. this bypasses local handle and did resolution, so its response decides which pds receives the app password. use an https endpoint; plain http is only accepted for localhost. + +## usage + +```text +usage: wisp-deploy [options] + +deploy a static site to wisp.place + +options: + -p, --path directory to deploy (required) + -s, --site site name (defaults to the directory name) + --directory enable directory listing + --spa serve index.html for routes without a file + -c, --concurrency number of concurrent uploads (default: 3) + --force-gzip gzip every file + --password app password for headless authentication + --password-file read the app password from a file + -h, --help show help ``` -## use in tangled ci +for example: + +```sh +export WISP_APP_PASSWORD='your app password' + +wisp-deploy \ + example.com \ + --path ./dist \ + --site example +``` -remember to set `WISP_APP_PASSWORD` in your spindle secrets +site names are at protocol record keys. they may contain letters, numbers, `.`, `-`, `_`, `:`, or `~`, but cannot be `.` or `..`. + +## files + +all regular files below `--path` are deployed unless they match the built-in ignore list or `.wispignore`. symlinks and other non-regular files are skipped. + +`.wispignore` uses gitignore-style patterns. the built-in list already skips version control directories, `node_modules`, virtual environments, caches, `.env` files, and `.wispignore` itself. point `--path` at the built site rather than the repository root unless you have checked everything that will be uploaded. + +wisp.place currently allows up to 1,000 files, 300 mib per site, and 200 mib per file. the cli stops before uploading when a deployment exceeds one of those limits. + +text files are gzipped before upload. unchanged blobs are reused from the existing manifest. uploads start with the requested concurrency and drop to two after a rate-limit response. + +## tangled ci + +add `WISP_APP_PASSWORD` to the spindle secrets, then use a step like this: ```yaml when: @@ -32,15 +68,15 @@ when: engine: "nixery" environment: - SITE_PATH: "." - SITE_NAME: "madoka" - WISP_HANDLE: "madoka.systems" + SITE_PATH: "dist" + SITE_NAME: "example" + WISP_HANDLE: "example.com" caches: - https://madoka-systems.cachix.org: "madoka-systems.cachix.org-1:nUYOriy5WXsFsO/eok8g/IgME2IT+6Vmna3itvuxJH8=" + https://madoka-systems.cachix.org: "madoka-systems.cachix.org-1:nUYOriy5WXsFsO/eok8g/IgME2IT+6Vmna3itvuxJH8=" dependencies: - - git+https://tangled.org/madoka.systems/wisp-deploy#default + - git+https://tangled.org/madoka.systems/wisp-deploy#default steps: - name: deploy to wisp @@ -48,10 +84,11 @@ steps: wisp-deploy \ "$WISP_HANDLE" \ --path "$SITE_PATH" \ - --site "$SITE_NAME" \ - --password "$WISP_APP_PASSWORD" + --site "$SITE_NAME" ``` -tags are also provided, you can use them by changing the url to `git+https://tangled.org/madoka.systems/wisp-deploy?ref=refs/tags/v1.0.0` for example. -for the sake of reproducibility, you can also use the url `git+https://tangled.org/did:plc:ne3mufdifemv72zui4bqlp32` instead if you desire. +use `?ref=refs/tags/v1.1.0` to pin this release. the did-based repository url is `git+https://tangled.org/did:plc:ne3mufdifemv72zui4bqlp32`. + +## license +MIT. see [`LICENSE`](LICENSE). diff --git a/flake.nix b/flake.nix index be8ebd6..5529076 100644 --- a/flake.nix +++ b/flake.nix @@ -27,7 +27,7 @@ packages = forAllSystems (pkgs: rec { wisp-deploy = pkgs.buildGoModule { pname = "wisp-deploy"; - version = "1.0.0"; + version = "1.1.0"; src = self; vendorHash = "sha256-WT7V95m8YjB2GdDoH11WIzBeh4YMUzIBLHaKbvXuxDg="; subPackages = [ "cmd/wisp-deploy" ]; diff --git a/internal/wisp/client.go b/internal/wisp/client.go index 0ca4683..5775cdb 100644 --- a/internal/wisp/client.go +++ b/internal/wisp/client.go @@ -3,11 +3,14 @@ package wisp import ( "bytes" "encoding/json" + "errors" "fmt" "io" + "net" "net/http" "net/url" "strings" + "time" ) type xrpcError struct { @@ -17,28 +20,69 @@ type xrpcError struct { type apiError struct { Status int + Code string Body string } func (e *apiError) Error() string { - return fmt.Sprintf("xrpc status %d: %s", e.Status, e.Body) + return fmt.Sprintf("xrpc status %d: %q", e.Status, e.Body) } // isStatus reports whether err is an *apiError with the given HTTP status. func isStatus(err error, code int) bool { - ae, ok := err.(*apiError) - return ok && ae.Status == code + var apiErr *apiError + return errors.As(err, &apiErr) && apiErr.Status == code +} +func isXRPCError(err error, code string) bool { + var apiErr *apiError + return errors.As(err, &apiErr) && apiErr.Code == code } type client struct { - pds string - jwt string - did string - hc *http.Client + pds string + jwt string + did string + hc *http.Client + sleep func(time.Duration) } -func newClient(pds string) *client { - return &client{pds: strings.TrimRight(pds, "/"), hc: http.DefaultClient} +func newHTTPClient() *http.Client { + transport := http.DefaultTransport.(*http.Transport).Clone() + transport.DialContext = (&net.Dialer{ + Timeout: 10 * time.Second, + KeepAlive: 30 * time.Second, + }).DialContext + transport.TLSHandshakeTimeout = 10 * time.Second + transport.ResponseHeaderTimeout = 30 * time.Second + transport.ExpectContinueTimeout = time.Second + + return &http.Client{ + Transport: transport, + Timeout: 5 * time.Minute, + CheckRedirect: func(req *http.Request, via []*http.Request) error { + if len(via) >= 10 { + return fmt.Errorf("stopped after 10 redirects") + } + if req.URL.Scheme != "https" || len(via) == 0 || !strings.EqualFold(req.URL.Host, via[0].URL.Host) { + return fmt.Errorf("refusing redirect to %s", req.URL.Redacted()) + } + return nil + }, + } +} + +func newClient(pds string, hc *http.Client) (*client, error) { + u, err := url.Parse(pds) + if err != nil { + return nil, fmt.Errorf("invalid PDS URL: %w", err) + } + if u.Scheme != "https" || u.Host == "" || u.User != nil || u.RawQuery != "" || u.Fragment != "" || (u.Path != "" && u.Path != "/") { + return nil, fmt.Errorf("invalid PDS URL %q: expected an HTTPS origin", pds) + } + if hc == nil { + hc = newHTTPClient() + } + return &client{pds: strings.TrimRight(u.String(), "/"), hc: hc, sleep: time.Sleep}, nil } func (c *client) do(method, nsid, contentType string, body []byte, query url.Values) ([]byte, error) { @@ -61,9 +105,32 @@ func (c *client) do(method, nsid, contentType string, body []byte, query url.Val return nil, err } defer resp.Body.Close() - data, _ := io.ReadAll(resp.Body) + + const maxResponseSize = 4 * 1024 * 1024 + const maxErrorSize = 64 * 1024 + limit := int64(maxResponseSize) + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + limit = maxErrorSize + } + data, err := readBounded(resp.Body, limit) + if err != nil { + return nil, err + } if resp.StatusCode < 200 || resp.StatusCode >= 300 { - return nil, &apiError{Status: resp.StatusCode, Body: string(data)} + var xrpc xrpcError + _ = json.Unmarshal(data, &xrpc) + return nil, &apiError{Status: resp.StatusCode, Code: xrpc.Error, Body: strings.TrimSpace(string(data))} + } + return data, nil +} + +func readBounded(r io.Reader, limit int64) ([]byte, error) { + data, err := io.ReadAll(io.LimitReader(r, limit+1)) + if err != nil { + return nil, err + } + if int64(len(data)) > limit { + return nil, fmt.Errorf("response exceeds %d bytes", limit) } return data, nil } @@ -97,7 +164,7 @@ func (c *client) getRecord(repo, collection, rkey string) (json.RawMessage, erro q.Set("rkey", rkey) data, err := c.do(http.MethodGet, "com.atproto.repo.getRecord", "", nil, q) if err != nil { - if isStatus(err, 400) || isStatus(err, 404) { + if isStatus(err, http.StatusNotFound) || isXRPCError(err, "RecordNotFound") { return nil, nil } return nil, err @@ -122,26 +189,35 @@ func (c *client) uploadBlob(content []byte) (json.RawMessage, error) { if err := json.Unmarshal(data, &out); err != nil { return nil, err } + if len(out.Blob) == 0 || string(out.Blob) == "null" { + return nil, fmt.Errorf("upload response missing blob") + } return out.Blob, nil } func (c *client) putRecord(repo, collection, rkey string, record any) error { - body, _ := json.Marshal(map[string]any{ + body, err := json.Marshal(map[string]any{ "repo": repo, "collection": collection, "rkey": rkey, "record": record, }) - _, err := c.do(http.MethodPost, "com.atproto.repo.putRecord", "application/json", body, nil) + if err != nil { + return err + } + _, err = c.do(http.MethodPost, "com.atproto.repo.putRecord", "application/json", body, nil) return err } func (c *client) deleteRecord(repo, collection, rkey string) error { - body, _ := json.Marshal(map[string]string{ + body, err := json.Marshal(map[string]string{ "repo": repo, "collection": collection, "rkey": rkey, }) - _, err := c.do(http.MethodPost, "com.atproto.repo.deleteRecord", "application/json", body, nil) + if err != nil { + return err + } + _, err = c.do(http.MethodPost, "com.atproto.repo.deleteRecord", "application/json", body, nil) return err } diff --git a/internal/wisp/collect_test.go b/internal/wisp/collect_test.go index 38325b1..abcb48f 100644 --- a/internal/wisp/collect_test.go +++ b/internal/wisp/collect_test.go @@ -53,3 +53,26 @@ func TestComputeCIDVectors(t *testing.T) { } } } + +func TestCollectFilesFailsWhenWispignoreIsUnreadable(t *testing.T) { + root := t.TempDir() + if err := os.Mkdir(filepath.Join(root, ".wispignore"), 0o755); err != nil { + t.Fatal(err) + } + if _, err := collectFiles(root); err == nil { + t.Fatal("expected .wispignore read error") + } +} + +func TestValidRecordKey(t *testing.T) { + for _, key := range []string{"site", "example.com", "a:b_c~d"} { + if !validRecordKey(key) { + t.Errorf("validRecordKey(%q) = false", key) + } + } + for _, key := range []string{"", ".", "..", "has/slash", strings.Repeat("a", 513)} { + if validRecordKey(key) { + t.Errorf("validRecordKey(%q) = true", key) + } + } +} diff --git a/internal/wisp/deploy.go b/internal/wisp/deploy.go index 9cd5f99..4bbb426 100644 --- a/internal/wisp/deploy.go +++ b/internal/wisp/deploy.go @@ -1,11 +1,12 @@ package wisp import ( + "errors" "flag" "fmt" "io" + "os" "path/filepath" - "regexp" "strings" "sync" ) @@ -16,27 +17,61 @@ const ( maxFileSize = 200 * 1024 * 1024 ) -var siteNameRe = regexp.MustCompile(`^[a-zA-Z0-9._~:-]{1,512}$`) - func printHelp() { - fmt.Println("Usage: wisp-deploy [options] ") + fmt.Println("usage: wisp-deploy [options] ") fmt.Println() - fmt.Println("Deploy a static site to wisp.place") + fmt.Println("deploy a static site to wisp.place") fmt.Println() - fmt.Println("Options:") + fmt.Println("options:") rows := [][2]string{ - {"-p, --path ", "Directory to deploy"}, - {"-s, --site ", "Site name (defaults to directory name)"}, - {"--directory", "Enable directory listing"}, - {"--spa", "Enable SPA mode (serve index.html for all routes)"}, - {"-c, --concurrency ", "Number of concurrent uploads (backs off to 2 on rate limit) (default: 3)"}, - {"--force-gzip", "Force gzip compression for all files regardless of type"}, - {"--password ", "App password for headless authentication"}, - {"-h, --help", "display help for command"}, - } - for _, r := range rows { - fmt.Printf(" %-21s %s\n", r[0], r[1]) + {"-p, --path ", "directory to deploy (required)"}, + {"-s, --site ", "site name (defaults to the directory name)"}, + {"--directory", "enable directory listing"}, + {"--spa", "serve index.html for routes without a file"}, + {"-c, --concurrency ", "number of concurrent uploads (default: 3)"}, + {"--force-gzip", "gzip every file"}, + {"--password ", "app password for headless authentication"}, + {"--password-file ", "read the app password from a file"}, + {"-h, --help", "show help"}, + } + for _, row := range rows { + fmt.Printf(" %-25s %s\n", row[0], row[1]) + } +} + +func loadPassword(flagValue, path string) (string, error) { + if flagValue != "" && path != "" { + return "", fmt.Errorf("use either --password or --password-file, not both") + } + if flagValue != "" { + password := strings.TrimSpace(flagValue) + if password == "" { + return "", fmt.Errorf("--password is empty") + } + return password, nil + } + if path == "" { + password := strings.TrimSpace(os.Getenv("WISP_APP_PASSWORD")) + if password == "" { + return "", fmt.Errorf("--password, --password-file, or WISP_APP_PASSWORD is required") + } + return password, nil + } + + f, err := os.Open(path) + if err != nil { + return "", fmt.Errorf("read password file: %w", err) } + defer f.Close() + data, err := readBounded(f, 64*1024) + if err != nil { + return "", fmt.Errorf("read password file: %w", err) + } + password := strings.TrimSpace(string(data)) + if password == "" { + return "", fmt.Errorf("password file is empty") + } + return password, nil } // Run executes the deploy CLI for the given args (excluding argv[0]). @@ -60,7 +95,8 @@ func Run(argv []string) error { fs.StringVar(path, "p", "", "directory to deploy (shorthand)") site := fs.String("site", "", "site name (defaults to directory name)") fs.StringVar(site, "s", "", "site name (shorthand)") - password := fs.String("password", "", "app password for headless authentication") + passwordFlag := fs.String("password", "", "app password for headless authentication") + passwordFile := fs.String("password-file", "", "read app password from a file") concurrency := fs.Int("concurrency", 3, "number of concurrent uploads") fs.IntVar(concurrency, "c", 3, "concurrent uploads (shorthand)") directory := fs.Bool("directory", false, "enable directory listing") @@ -80,31 +116,42 @@ func Run(argv []string) error { if handle == "" { return fmt.Errorf("handle is required") } - if *password == "" { - return fmt.Errorf("--password is required") - } if *path == "" { return fmt.Errorf("--path is required") } + if *concurrency < 1 || *concurrency > maxUploadConcurrency { + return fmt.Errorf("--concurrency must be between 1 and %d", maxUploadConcurrency) + } + password, err := loadPassword(*passwordFlag, *passwordFile) + if err != nil { + return err + } siteName := strings.ToLower(*site) if siteName == "" { siteName = strings.ToLower(filepath.Base(*path)) } - if !siteNameRe.MatchString(siteName) { - return fmt.Errorf("invalid site name %q: must be 1-512 chars of [a-zA-Z0-9._~:-]", siteName) + if !validRecordKey(siteName) { + return fmt.Errorf("invalid site name %q: expected an AT Protocol record key", siteName) } - doc, err := resolveIdentity(handle) + hc := newHTTPClient() + doc, err := resolveIdentity(hc, handle) if err != nil { return err } - fmt.Printf("Resolved %s -> %s (%s)\n", handle, doc.DID, doc.PDS) - c := newClient(doc.PDS) - if err := c.createSession(handle, strings.TrimSpace(*password)); err != nil { + c, err := newClient(doc.PDS, hc) + if err != nil { + return err + } + fmt.Printf("Resolved %s -> %s (%s)\n", handle, doc.DID, doc.PDS) + if err := c.createSession(handle, password); err != nil { return fmt.Errorf("login failed: %w", err) } + if c.did != doc.DID { + return fmt.Errorf("login returned DID %s, expected %s", c.did, doc.DID) + } fmt.Printf("Authenticated as %s\n", c.did) return deploy(c, siteName, *path, *concurrency, *forceGzip, *directory, *spa, handle) @@ -123,18 +170,18 @@ func deploy(c *client, site, dir string, concurrency int, forceGzip, directory, for _, f := range files { total += f.size if f.size > maxFileSize { - fmt.Printf("Warning: %s exceeds max file size and may not be cached\n", f.rel) + return fmt.Errorf("%s is %s; maximum file size is %s", f.rel, formatBytes(f.size), formatBytes(maxFileSize)) } } fmt.Printf("Found %d files (%s)\n", len(files), formatBytes(total)) if len(files) > maxFileCount { - fmt.Printf("Warning: %d files exceeds limit %d; site may not be cached\n", len(files), maxFileCount) + return fmt.Errorf("site has %d files; maximum is %d", len(files), maxFileCount) } if total > maxSiteSize { - fmt.Printf("Warning: site size %s exceeds limit; site may not be cached\n", formatBytes(total)) + return fmt.Errorf("site is %s; maximum size is %s", formatBytes(total), formatBytes(maxSiteSize)) } - reuse, oldRoot, err := c.fetchExisting(site) + reuse, oldSubfsRkeys, err := c.fetchExisting(site) if err != nil { return err } @@ -168,8 +215,14 @@ func deploy(c *client, site, dir string, concurrency int, forceGzip, directory, } finalRoot, rkeys = fr, rk } - manifest := FsRecord{Type: collFs, Site: site, Root: finalRoot, FileCount: countFiles(finalRoot), CreatedAt: timeNowISO()} - return rkeys, c.putRecord(c.did, collFs, site, manifest) + manifest := FsRecord{Type: collFs, Site: site, Root: finalRoot, FileCount: fileCount, CreatedAt: timeNowISO()} + if err := c.putRecord(c.did, collFs, site, manifest); err != nil { + if cleanupErr := c.deleteSubfsRecords(rkeys); cleanupErr != nil { + return nil, errors.Join(err, cleanupErr) + } + return nil, err + } + return rkeys, nil } subfsRkeys, err := putManifest() @@ -194,20 +247,19 @@ func deploy(c *client, site, dir string, concurrency int, forceGzip, directory, } } - if oldRoot != nil { - newSet := map[string]bool{} - for _, k := range subfsRkeys { - newSet[k] = true - } - var oldURIs []string - collectSubfsURIs(oldRoot, &oldURIs) - for _, uri := range oldURIs { - _, rkey, ok := parseSubfsURI(uri) - if ok && !newSet[rkey] { - _ = c.deleteRecord(c.did, collSubfs, rkey) - } + newSet := make(map[string]bool, len(subfsRkeys)) + for _, rkey := range subfsRkeys { + newSet[rkey] = true + } + var obsolete []string + for _, rkey := range oldSubfsRkeys { + if !newSet[rkey] { + obsolete = append(obsolete, rkey) } } + if err := c.deleteSubfsRecords(obsolete); err != nil { + fmt.Printf("Warning: could not remove every old subfs record: %v\n", err) + } if directory || spa { settings := map[string]any{ diff --git a/internal/wisp/deploy_test.go b/internal/wisp/deploy_test.go new file mode 100644 index 0000000..41ef193 --- /dev/null +++ b/internal/wisp/deploy_test.go @@ -0,0 +1,111 @@ +package wisp + +import ( + "encoding/json" + "io" + "net/http" + "os" + "path/filepath" + "sync" + "testing" + "time" +) + +type publicationRT struct { + mu sync.Mutex + created []string + deleted []string + manifestPuts int + manifestCount int +} + +func (rt *publicationRT) RoundTrip(req *http.Request) (*http.Response, error) { + nsid := filepath.Base(req.URL.Path) + switch nsid { + case "com.atproto.repo.getRecord": + return testResponse(http.StatusNotFound, `{"error":"RecordNotFound"}`), nil + case "com.atproto.repo.uploadBlob": + body, err := io.ReadAll(req.Body) + if err != nil { + return nil, err + } + response := `{"blob":{"ref":{"$link":"` + computeCID(body) + `"}}}` + return testResponse(http.StatusOK, response), nil + case "com.atproto.repo.putRecord": + var call putCall + if err := json.NewDecoder(req.Body).Decode(&call); err != nil { + return nil, err + } + rt.mu.Lock() + defer rt.mu.Unlock() + if call.Collection == collSubfs { + rt.created = append(rt.created, call.Rkey) + return testResponse(http.StatusOK, `{}`), nil + } + if call.Collection == collFs { + rt.manifestPuts++ + var record FsRecord + if err := json.Unmarshal(call.Record, &record); err != nil { + return nil, err + } + rt.manifestCount = record.FileCount + return testResponse(http.StatusInternalServerError, `{"error":"InternalError"}`), nil + } + case "com.atproto.repo.deleteRecord": + var call putCall + if err := json.NewDecoder(req.Body).Decode(&call); err != nil { + return nil, err + } + rt.mu.Lock() + rt.deleted = append(rt.deleted, call.Rkey) + rt.mu.Unlock() + return testResponse(http.StatusOK, `{}`), nil + } + return testResponse(http.StatusNotFound, `{}`), nil +} + +func TestFailedManifestPublicationCleansNewSubfsGeneration(t *testing.T) { + root := t.TempDir() + for i := range 250 { + name := filepath.Join(root, "file-"+formatIndex(i)+".bin") + if err := os.WriteFile(name, []byte{byte(i)}, 0o644); err != nil { + t.Fatal(err) + } + } + + rt := &publicationRT{} + c := testClient(rt) + c.sleep = func(time.Duration) {} + if err := deploy(c, "site", root, 3, false, false, false, "example.com"); err == nil { + t.Fatal("expected manifest publication to fail") + } + + rt.mu.Lock() + defer rt.mu.Unlock() + if rt.manifestPuts != 2 { + t.Fatalf("manifest puts=%d, want 2", rt.manifestPuts) + } + if rt.manifestCount != 250 { + t.Fatalf("manifest fileCount=%d, want 250", rt.manifestCount) + } + if len(rt.created) == 0 || len(rt.created) != len(rt.deleted) { + t.Fatalf("created=%v deleted=%v", rt.created, rt.deleted) + } + seen := make(map[string]bool, len(rt.created)) + for _, key := range rt.created { + if seen[key] { + t.Fatalf("subfs key reused across attempts: %s", key) + } + seen[key] = true + } + for _, key := range rt.deleted { + if !seen[key] { + t.Fatalf("deleted unknown key %s", key) + } + } +} + +func formatIndex(value int) string { + const digits = "0123456789" + return string([]byte{digits[value/100], digits[value/10%10], digits[value%10]}) +} diff --git a/internal/wisp/identity.go b/internal/wisp/identity.go index f1f9d57..b84a83f 100644 --- a/internal/wisp/identity.go +++ b/internal/wisp/identity.go @@ -1,51 +1,247 @@ package wisp import ( + "context" "encoding/json" "fmt" + "net" "net/http" "net/url" "os" + "strings" + "time" ) -const defaultMiniDocURL = "https://slingshot.microcosm.blue/xrpc/blue.microcosm.identity.resolveMiniDoc" - type miniDoc struct { DID string `json:"did"` Handle string `json:"handle"` PDS string `json:"pds"` } -// resolveIdentity turns a handle or DID into its DID and PDS endpoint in a -// single request via slingshot's resolveMiniDoc. -func resolveIdentity(identifier string) (miniDoc, error) { - base := os.Getenv("WISP_MINIDOC_URL") - if base == "" { - base = defaultMiniDocURL +type didDocument struct { + ID string `json:"id"` + Service []didService `json:"service"` +} + +type didService struct { + ID string `json:"id"` + Type json.RawMessage `json:"type"` + ServiceEndpoint string `json:"serviceEndpoint"` +} + +func resolveIdentity(hc *http.Client, identifier string) (miniDoc, error) { + identifier = strings.TrimSpace(identifier) + if identifier == "" { + return miniDoc{}, fmt.Errorf("identifier is empty") + } + if endpoint := strings.TrimSpace(os.Getenv("WISP_MINIDOC_URL")); endpoint != "" { + return resolveMiniDoc(hc, endpoint, identifier) + } + + did := identifier + handle := "" + if !strings.HasPrefix(identifier, "did:") { + var err error + handle, err = normalizeHandle(identifier) + if err != nil { + return miniDoc{}, err + } + did, err = resolveHandle(hc, handle) + if err != nil { + return miniDoc{}, err + } + } + + doc, err := resolveDIDDocument(hc, did) + if err != nil { + return miniDoc{}, err + } + if doc.ID != did { + return miniDoc{}, fmt.Errorf("DID document id %q does not match %q", doc.ID, did) + } + for _, service := range doc.Service { + if strings.HasSuffix(service.ID, "#atproto_pds") && serviceHasType(service.Type, "AtprotoPersonalDataServer") { + if service.ServiceEndpoint == "" { + break + } + return miniDoc{DID: did, Handle: handle, PDS: service.ServiceEndpoint}, nil + } + } + return miniDoc{}, fmt.Errorf("DID document for %s has no AT Protocol PDS service", did) +} +func resolveMiniDoc(hc *http.Client, endpoint, identifier string) (miniDoc, error) { + u, err := url.Parse(endpoint) + if err != nil { + return miniDoc{}, fmt.Errorf("invalid WISP_MINIDOC_URL: %w", err) + } + isLocalHTTP := u.Scheme == "http" && (u.Hostname() == "localhost" || net.ParseIP(u.Hostname()).IsLoopback()) + if (u.Scheme != "https" && !isLocalHTTP) || u.Host == "" || u.User != nil || u.Fragment != "" { + return miniDoc{}, fmt.Errorf("invalid WISP_MINIDOC_URL %q: expected HTTPS or local HTTP", endpoint) } - u := fmt.Sprintf("%s?identifier=%s", base, url.QueryEscape(identifier)) + query := u.Query() + query.Set("identifier", identifier) + u.RawQuery = query.Encode() - resp, err := http.Get(u) + resp, err := hc.Get(u.String()) if err != nil { return miniDoc{}, fmt.Errorf("resolve %s: %w", identifier, err) } defer resp.Body.Close() - + data, err := readBounded(resp.Body, 1024*1024) + if err != nil { + return miniDoc{}, fmt.Errorf("resolve %s: %w", identifier, err) + } if resp.StatusCode != http.StatusOK { - var e xrpcError - _ = json.NewDecoder(resp.Body).Decode(&e) - if e.Message != "" { - return miniDoc{}, fmt.Errorf("resolve %s: %s", identifier, e.Message) - } - return miniDoc{}, fmt.Errorf("resolve %s: status %d", identifier, resp.StatusCode) + return miniDoc{}, fmt.Errorf("resolve %s: status %d: %q", identifier, resp.StatusCode, strings.TrimSpace(string(data))) } - var doc miniDoc - if err := json.NewDecoder(resp.Body).Decode(&doc); err != nil { + if err := json.Unmarshal(data, &doc); err != nil { return miniDoc{}, fmt.Errorf("decode minidoc: %w", err) } - if doc.DID == "" || doc.PDS == "" { - return miniDoc{}, fmt.Errorf("incomplete identity for %s", identifier) + did, err := normalizeDID(doc.DID) + if err != nil { + return miniDoc{}, fmt.Errorf("invalid minidoc DID: %w", err) + } + if doc.PDS == "" { + return miniDoc{}, fmt.Errorf("minidoc for %s has no PDS", identifier) } + doc.DID = did return doc, nil } + +func normalizeHandle(handle string) (string, error) { + handle = strings.ToLower(strings.TrimSuffix(strings.TrimSpace(handle), ".")) + if len(handle) < 3 || len(handle) > 253 { + return "", fmt.Errorf("invalid handle %q", handle) + } + labels := strings.Split(handle, ".") + if len(labels) < 2 { + return "", fmt.Errorf("invalid handle %q", handle) + } + for _, label := range labels { + if len(label) == 0 || len(label) > 63 || label[0] == '-' || label[len(label)-1] == '-' { + return "", fmt.Errorf("invalid handle %q", handle) + } + for _, ch := range label { + if (ch < 'a' || ch > 'z') && (ch < '0' || ch > '9') && ch != '-' { + return "", fmt.Errorf("invalid handle %q", handle) + } + } + } + return handle, nil +} + +func resolveHandle(hc *http.Client, handle string) (string, error) { + wellKnown := "https://" + handle + "/.well-known/atproto-did" + resp, httpsErr := hc.Get(wellKnown) + if httpsErr == nil { + data, readErr := readBounded(resp.Body, 8*1024) + resp.Body.Close() + switch { + case resp.StatusCode != http.StatusOK: + httpsErr = fmt.Errorf("status %d", resp.StatusCode) + case readErr != nil: + httpsErr = readErr + default: + did, err := normalizeDID(string(data)) + if err == nil { + return did, nil + } + httpsErr = err + } + } + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + records, dnsErr := net.DefaultResolver.LookupTXT(ctx, "_atproto."+handle) + if dnsErr == nil { + for _, record := range records { + if strings.HasPrefix(record, "did=") { + if did, err := normalizeDID(strings.TrimPrefix(record, "did=")); err == nil { + return did, nil + } + } + } + dnsErr = fmt.Errorf("no valid did= TXT record") + } + return "", fmt.Errorf("resolve handle %s: HTTPS lookup failed: %v; DNS lookup failed: %w", handle, httpsErr, dnsErr) +} + +func normalizeDID(value string) (string, error) { + did := strings.TrimSpace(value) + if strings.ContainsAny(did, " \t\r\n/?#") { + return "", fmt.Errorf("unsupported DID %q", did) + } + switch { + case strings.HasPrefix(did, "did:plc:"): + suffix := strings.TrimPrefix(did, "did:plc:") + if len(suffix) != 24 { + return "", fmt.Errorf("invalid did:plc identifier %q", did) + } + for _, ch := range suffix { + if (ch < 'a' || ch > 'z') && (ch < '2' || ch > '7') { + return "", fmt.Errorf("invalid did:plc identifier %q", did) + } + } + case strings.HasPrefix(did, "did:web:"): + if _, err := normalizeHandle(strings.TrimPrefix(did, "did:web:")); err != nil { + return "", fmt.Errorf("invalid did:web identifier %q", did) + } + default: + return "", fmt.Errorf("unsupported DID %q", did) + } + return did, nil +} + +func resolveDIDDocument(hc *http.Client, did string) (didDocument, error) { + did, err := normalizeDID(did) + if err != nil { + return didDocument{}, err + } + + var endpoint string + switch { + case strings.HasPrefix(did, "did:plc:"): + endpoint = "https://plc.directory/" + url.PathEscape(did) + case strings.HasPrefix(did, "did:web:"): + host, err := normalizeHandle(strings.TrimPrefix(did, "did:web:")) + if err != nil { + return didDocument{}, fmt.Errorf("invalid did:web identifier: %w", err) + } + endpoint = "https://" + host + "/.well-known/did.json" + } + + resp, err := hc.Get(endpoint) + if err != nil { + return didDocument{}, fmt.Errorf("resolve %s: %w", did, err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return didDocument{}, fmt.Errorf("resolve %s: status %d", did, resp.StatusCode) + } + data, err := readBounded(resp.Body, 1024*1024) + if err != nil { + return didDocument{}, fmt.Errorf("resolve %s: %w", did, err) + } + var doc didDocument + if err := json.Unmarshal(data, &doc); err != nil { + return didDocument{}, fmt.Errorf("decode DID document: %w", err) + } + return doc, nil +} + +func serviceHasType(raw json.RawMessage, want string) bool { + var single string + if json.Unmarshal(raw, &single) == nil { + return single == want + } + var many []string + if json.Unmarshal(raw, &many) == nil { + for _, value := range many { + if value == want { + return true + } + } + } + return false +} diff --git a/internal/wisp/security_test.go b/internal/wisp/security_test.go new file mode 100644 index 0000000..e739340 --- /dev/null +++ b/internal/wisp/security_test.go @@ -0,0 +1,114 @@ +package wisp + +import ( + "io" + "net/http" + "os" + "path/filepath" + "strings" + "testing" +) + +type roundTripFunc func(*http.Request) (*http.Response, error) + +func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) { + return f(req) +} + +func testResponse(status int, body string) *http.Response { + return &http.Response{ + StatusCode: status, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(body)), + } +} + +func TestNewClientRequiresHTTPSOrigin(t *testing.T) { + for _, endpoint := range []string{ + "http://pds.test", + "https://user@pds.test", + "https://pds.test/xrpc", + "https://pds.test?query=yes", + } { + if _, err := newClient(endpoint, &http.Client{}); err == nil { + t.Errorf("newClient(%q) accepted an unsafe endpoint", endpoint) + } + } + if _, err := newClient("https://pds.test", &http.Client{}); err != nil { + t.Fatal(err) + } +} + +func TestReadBoundedRejectsOversizedBody(t *testing.T) { + if _, err := readBounded(strings.NewReader("123"), 2); err == nil { + t.Fatal("expected an oversized response error") + } +} + +func TestResolveIdentityFromHandle(t *testing.T) { + t.Setenv("WISP_MINIDOC_URL", "") + const did = "did:plc:abcdefghijklmnopqrstuvwx" + hc := &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + switch req.URL.String() { + case "https://alice.test/.well-known/atproto-did": + return testResponse(http.StatusOK, did+"\n"), nil + case "https://plc.directory/" + did: + return testResponse(http.StatusOK, `{"id":"`+did+`","service":[{"id":"`+did+`#atproto_pds","type":"AtprotoPersonalDataServer","serviceEndpoint":"https://pds.test"}]}`), nil + default: + t.Fatalf("unexpected request: %s", req.URL) + return nil, nil + } + })} + + doc, err := resolveIdentity(hc, "Alice.Test") + if err != nil { + t.Fatal(err) + } + if doc.DID != did || doc.Handle != "alice.test" || doc.PDS != "https://pds.test" { + t.Fatalf("resolved %#v", doc) + } +} +func TestResolveIdentityFromMinidocOverride(t *testing.T) { + const did = "did:plc:abcdefghijklmnopqrstuvwx" + t.Setenv("WISP_MINIDOC_URL", "https://slingshot.test/xrpc/blue.microcosm.identity.resolveMiniDoc") + hc := &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + if req.URL.Query().Get("identifier") != "alice.test" { + t.Fatalf("identifier=%q", req.URL.Query().Get("identifier")) + } + return testResponse(http.StatusOK, `{"did":"`+did+`","handle":"alice.test","pds":"https://pds.test"}`), nil + })} + + doc, err := resolveIdentity(hc, "alice.test") + if err != nil { + t.Fatal(err) + } + if doc.DID != did || doc.PDS != "https://pds.test" { + t.Fatalf("resolved %#v", doc) + } +} + +func TestLoadPassword(t *testing.T) { + t.Setenv("WISP_APP_PASSWORD", " env-secret \n") + password, err := loadPassword("", "") + if err != nil || password != "env-secret" { + t.Fatalf("password=%q err=%v", password, err) + } + password, err = loadPassword("flag-secret", "") + if err != nil || password != "flag-secret" { + t.Fatalf("password=%q err=%v", password, err) + } + + path := filepath.Join(t.TempDir(), "password") + if err := os.WriteFile(path, []byte("file-secret\n"), 0o600); err != nil { + t.Fatal(err) + } + t.Setenv("WISP_APP_PASSWORD", "") + password, err = loadPassword("", path) + if err != nil || password != "file-secret" { + t.Fatalf("password=%q err=%v", password, err) + } + + if _, err := loadPassword("flag-secret", path); err == nil { + t.Fatal("expected conflicting password flags to fail") + } +} diff --git a/internal/wisp/split_test.go b/internal/wisp/split_test.go index 2210d01..b0a6cdd 100644 --- a/internal/wisp/split_test.go +++ b/internal/wisp/split_test.go @@ -40,8 +40,10 @@ func (m *recordingRT) RoundTrip(req *http.Request) (*http.Response, error) { } func testClient(rt http.RoundTripper) *client { - c := newClient("https://pds.test") - c.hc = &http.Client{Transport: rt} + c, err := newClient("https://pds.test", &http.Client{Transport: rt}) + if err != nil { + panic(err) + } c.did = "did:plc:test" c.jwt = "token" return c @@ -119,16 +121,16 @@ func TestSplitIntoSubfsSingleDir(t *testing.T) { if err != nil { t.Fatal(err) } - if len(rkeys) != 1 || rkeys[0] != "mysite-subfs-1" { - t.Fatalf("rkeys = %v, want [mysite-subfs-1]", rkeys) + if len(rkeys) != 1 || !validRecordKey(rkeys[0]) { + t.Fatalf("invalid generated rkeys: %v", rkeys) } if len(rt.calls) != 1 { t.Fatalf("putRecord calls = %d, want 1", len(rt.calls)) } call := rt.calls[0] - if call.Collection != collSubfs || call.Rkey != "mysite-subfs-1" { - t.Fatalf("subfs put to %s/%s", call.Collection, call.Rkey) + if call.Collection != collSubfs || call.Rkey != rkeys[0] { + t.Fatalf("subfs put to %s/%s, expected %s", call.Collection, call.Rkey, rkeys[0]) } var rec SubfsRecord if err := json.Unmarshal(call.Record, &rec); err != nil { @@ -146,7 +148,8 @@ func TestSplitIntoSubfsSingleDir(t *testing.T) { t.Fatalf("final root entries = %d", len(finalRoot.Entries)) } node := finalRoot.Entries[0].Node - if node.NodeType != "subfs" || node.Subject != "at://did:plc:test/place.wisp.subfs/mysite-subfs-1" { + wantSubject := "at://did:plc:test/place.wisp.subfs/" + rkeys[0] + if node.NodeType != "subfs" || node.Subject != wantSubject { t.Fatalf("expected subfs node, got type=%s subject=%s", node.NodeType, node.Subject) } if node.Flat == nil || *node.Flat != false { @@ -174,18 +177,20 @@ func TestSplitIntoSubfsChunks(t *testing.T) { for _, call := range rt.calls { byRkey[call.Rkey] = call } - for _, k := range rkeys { - if strings.Contains(k, "-chunk-") { - chunks += k + " " - } else { - parent = k + for _, key := range rkeys { + kind := key[strings.LastIndexByte(key, '-')+1:] + switch { + case strings.HasPrefix(kind, "c"): + chunks += key + " " + case strings.HasPrefix(kind, "p"): + parent = key } } if chunks == "" { t.Fatalf("expected chunk records, got rkeys %v", rkeys) } - if parent != "mysite-subfs-1" { - t.Fatalf("parent rkey = %q, want mysite-subfs-1", parent) + if parent == "" { + t.Fatalf("expected parent record, got rkeys %v", rkeys) } parentBody := string(byRkey[parent].Record) @@ -196,3 +201,36 @@ func TestSplitIntoSubfsChunks(t *testing.T) { t.Errorf("subfs#subfs chunk references must not carry flat\ngot: %s", parentBody) } } + +func TestSubfsKeysAreGenerationSpecificAndBounded(t *testing.T) { + var firstNonce, secondNonce [8]byte + firstNonce[7] = 1 + secondNonce[7] = 2 + first, err := subfsKeyGeneratorFromNonce(strings.Repeat("a", 512), firstNonce).key("d") + if err != nil { + t.Fatal(err) + } + second, err := subfsKeyGeneratorFromNonce(strings.Repeat("a", 512), secondNonce).key("d") + if err != nil { + t.Fatal(err) + } + if first == second { + t.Fatal("different generations produced the same key") + } + if !validRecordKey(first) || !validRecordKey(second) { + t.Fatalf("invalid keys %q %q", first, second) + } +} + +func TestSplitIntoSubfsHandlesDeepLargeDirectory(t *testing.T) { + rt := &recordingRT{} + c := testClient(rt) + root := newFsDir() + outer := newFsDir() + outer.Entries = append(outer.Entries, &FsEntry{Name: "inner", Node: dirWithFiles(800)}) + root.Entries = append(root.Entries, &FsEntry{Name: "outer", Node: outer}) + + if _, _, err := c.splitIntoSubfs(root, "mysite"); err != nil { + t.Fatal(err) + } +} diff --git a/internal/wisp/subfs.go b/internal/wisp/subfs.go index caad5d7..aad994d 100644 --- a/internal/wisp/subfs.go +++ b/internal/wisp/subfs.go @@ -1,7 +1,11 @@ package wisp import ( + "crypto/rand" + "crypto/sha256" + "encoding/hex" "encoding/json" + "errors" "fmt" "strings" ) @@ -68,7 +72,7 @@ func replaceDirWithSubfs(root *FsNode, targetPath, uri string, flat bool) bool { return false } -func splitDirIntoChunks(dir *FsNode, maxSize int) []*FsNode { +func splitDirIntoChunks(dir *FsNode, maxSize int) ([]*FsNode, error) { var chunks []*FsNode var cur []*FsEntry curSize := 100 @@ -81,17 +85,23 @@ func splitDirIntoChunks(dir *FsNode, maxSize int) []*FsNode { curSize = 100 } } - for _, e := range dir.Entries { - b, _ := json.Marshal(e) - es := len(b) - if len(cur) > 0 && curSize+es > maxSize { + for _, entry := range dir.Entries { + body, err := json.Marshal(entry) + if err != nil { + return nil, err + } + entrySize := len(body) + if entrySize+100 > maxSize { + return nil, fmt.Errorf("entry %q is too large for a subfs chunk", entry.Name) + } + if len(cur) > 0 && curSize+entrySize > maxSize { flush() } - cur = append(cur, e) - curSize += es + cur = append(cur, entry) + curSize += entrySize } flush() - return chunks + return chunks, nil } func (c *client) createSubfsRecord(dir *FsNode, rkey string) (string, error) { @@ -102,41 +112,113 @@ func (c *client) createSubfsRecord(dir *FsNode, rkey string) (string, error) { return fmt.Sprintf("at://%s/%s/%s", c.did, collSubfs, rkey), nil } +type subfsKeyGenerator struct { + prefix string + next int +} + +func newSubfsKeyGenerator(site string) (*subfsKeyGenerator, error) { + var nonce [8]byte + if _, err := rand.Read(nonce[:]); err != nil { + return nil, fmt.Errorf("generate subfs key: %w", err) + } + return subfsKeyGeneratorFromNonce(site, nonce), nil +} + +func subfsKeyGeneratorFromNonce(site string, nonce [8]byte) *subfsKeyGenerator { + siteHash := sha256.Sum256([]byte(site)) + return &subfsKeyGenerator{ + prefix: "wd-" + hex.EncodeToString(siteHash[:8]) + "-" + hex.EncodeToString(nonce[:]), + } +} + +func (g *subfsKeyGenerator) key(kind string) (string, error) { + g.next++ + key := fmt.Sprintf("%s-%s%d", g.prefix, kind, g.next) + if !validRecordKey(key) { + return "", fmt.Errorf("generated invalid subfs key %q", key) + } + return key, nil +} + +func (c *client) deleteSubfsRecords(rkeys []string) error { + var errs []error + for i := len(rkeys) - 1; i >= 0; i-- { + if err := c.deleteRecord(c.did, collSubfs, rkeys[i]); err != nil && !isStatus(err, 404) { + errs = append(errs, fmt.Errorf("delete subfs %s: %w", rkeys[i], err)) + } + } + return errors.Join(errs...) +} + +func canChunkDirectory(dir *FsNode, maxSize int) bool { + for _, entry := range dir.Entries { + body, err := json.Marshal(entry) + if err != nil || len(body)+100 > maxSize { + return false + } + } + return true +} + // splitIntoSubfs reduces the manifest below the PDS record-size limit by -// extracting large directories into separate place.wisp.subfs records. -func (c *client) splitIntoSubfs(root *FsNode, site string) (*FsNode, []string, error) { - var subfsRkeys []string +// extracting large directories into generation-specific subfs records. +func (c *client) splitIntoSubfs(root *FsNode, site string) (finalRoot *FsNode, rkeys []string, resultErr error) { + keys, err := newSubfsKeyGenerator(site) + if err != nil { + return nil, nil, err + } + var created []string + complete := false + defer func() { + if !complete { + if cleanupErr := c.deleteSubfsRecords(created); cleanupErr != nil { + resultErr = errors.Join(resultErr, cleanupErr) + } + } + }() + writeRecord := func(dir *FsNode, kind string) (string, error) { + rkey, err := keys.key(kind) + if err != nil { + return "", err + } + created = append(created, rkey) + return c.createSubfsRecord(dir, rkey) + } + cur := root curFileCount := countFiles(cur) - chunkCounter := 0 - - for iter := 1; iter <= 100; iter++ { + for range 100 { if manifestSize(site, cur, curFileCount) <= maxManifestSize && curFileCount <= targetFileCount { - break + complete = true + return cur, created, nil } dirs := findDirectories(cur, "") - if len(dirs) > 0 { - largest := dirs[0] - for _, d := range dirs { - if d.size > largest.size { - largest = d - } + var largest *dirRef + for i := range dirs { + dir := &dirs[i] + if dir.size > maxSubfsSize && !canChunkDirectory(dir.dir, maxSubfsSize) { + continue } - + if largest == nil || dir.size > largest.size { + largest = dir + } + } + if largest != nil { var subfsURI string if largest.size > maxSubfsSize { - chunks := splitDirIntoChunks(largest.dir, maxSubfsSize) + chunks, err := splitDirIntoChunks(largest.dir, maxSubfsSize) + if err != nil { + return nil, nil, err + } var chunkURIs []string for _, chunk := range chunks { - rkey := fmt.Sprintf("%s-chunk-%d", site, chunkCounter) - chunkCounter++ - uri, err := c.createSubfsRecord(chunk, rkey) + uri, err := writeRecord(chunk, "c") if err != nil { return nil, nil, err } chunkURIs = append(chunkURIs, uri) - subfsRkeys = append(subfsRkeys, rkey) } parent := newFsDir() for i, uri := range chunkURIs { @@ -145,20 +227,16 @@ func (c *client) splitIntoSubfs(root *FsNode, site string) (*FsNode, []string, e Node: newFsSubfs(uri, true), }) } - parentRkey := fmt.Sprintf("%s-subfs-%d", site, iter) - uri, err := c.createSubfsRecord(parent, parentRkey) + uri, err := writeRecord(parent, "p") if err != nil { return nil, nil, err } - subfsRkeys = append(subfsRkeys, parentRkey) subfsURI = uri } else { - rkey := fmt.Sprintf("%s-subfs-%d", site, iter) - uri, err := c.createSubfsRecord(largest.dir, rkey) + uri, err := writeRecord(largest.dir, "d") if err != nil { return nil, nil, err } - subfsRkeys = append(subfsRkeys, rkey) subfsURI = uri } @@ -169,43 +247,37 @@ func (c *client) splitIntoSubfs(root *FsNode, site string) (*FsNode, []string, e continue } - // No subdirectories: split flat root files into a subfs chunk. var rootFiles []*FsEntry - for _, e := range cur.Entries { - if e.Node.NodeType == "file" { - rootFiles = append(rootFiles, e) + for _, entry := range cur.Entries { + if entry.Node.NodeType == "file" { + rootFiles = append(rootFiles, entry) } } if len(rootFiles) == 0 { return nil, nil, fmt.Errorf("cannot split manifest further (%d files, %d bytes)", curFileCount, manifestSize(site, cur, curFileCount)) } const chunkSize = 100 - take := chunkSize - if take > len(rootFiles) { - take = len(rootFiles) - } + take := min(chunkSize, len(rootFiles)) chunkFiles := rootFiles[:take] chunkDir := newFsDir() chunkDir.Entries = chunkFiles - rkey := fmt.Sprintf("%s-subfs-%d", site, iter) - uri, err := c.createSubfsRecord(chunkDir, rkey) + uri, err := writeRecord(chunkDir, "f") if err != nil { return nil, nil, err } - subfsRkeys = append(subfsRkeys, rkey) taken := make(map[*FsEntry]bool, len(chunkFiles)) - for _, e := range chunkFiles { - taken[e] = true + for _, entry := range chunkFiles { + taken[entry] = true } var remaining []*FsEntry - for _, e := range cur.Entries { - if !taken[e] { - remaining = append(remaining, e) + for _, entry := range cur.Entries { + if !taken[entry] { + remaining = append(remaining, entry) } } remaining = append(remaining, &FsEntry{ - Name: fmt.Sprintf("__subfs_%d", iter), + Name: fmt.Sprintf("__subfs_%d", len(created)), Node: newFsSubfs(uri, true), }) next := newFsDir() @@ -214,5 +286,5 @@ func (c *client) splitIntoSubfs(root *FsNode, site string) (*FsNode, []string, e curFileCount -= take } - return cur, subfsRkeys, nil + return nil, nil, fmt.Errorf("cannot split manifest below limits after 100 iterations") } diff --git a/internal/wisp/upload.go b/internal/wisp/upload.go index e9017b7..c99052c 100644 --- a/internal/wisp/upload.go +++ b/internal/wisp/upload.go @@ -3,13 +3,16 @@ package wisp import ( "encoding/base64" "encoding/json" + "errors" "fmt" "io/fs" + "net/http" "os" "path/filepath" "sort" "strings" "sync" + "sync/atomic" "time" gitignore "github.com/sabhiram/go-gitignore" @@ -34,7 +37,11 @@ type fileInfo struct { func collectFiles(siteDir string) ([]fileInfo, error) { lines := append([]string{}, defaultIgnore...) - if extra, err := os.ReadFile(filepath.Join(siteDir, ".wispignore")); err == nil { + extra, err := os.ReadFile(filepath.Join(siteDir, ".wispignore")) + if err != nil && !errors.Is(err, os.ErrNotExist) { + return nil, fmt.Errorf("read .wispignore: %w", err) + } + if err == nil { for _, l := range strings.Split(string(extra), "\n") { l = strings.TrimSpace(l) if l != "" && !strings.HasPrefix(l, "#") { @@ -45,7 +52,7 @@ func collectFiles(siteDir string) ([]fileInfo, error) { ig := gitignore.CompileIgnoreLines(lines...) var files []fileInfo - err := filepath.WalkDir(siteDir, func(p string, d fs.DirEntry, err error) error { + err = filepath.WalkDir(siteDir, func(p string, d fs.DirEntry, err error) error { if err != nil { return err } @@ -82,45 +89,102 @@ func collectFiles(siteDir string) ([]fileInfo, error) { return files, nil } -// fetchExisting builds the path->blob reuse map from an existing fs record and -// its referenced subfs records. -func (c *client) fetchExisting(site string) (map[string]blobInfo, *FsNode, error) { +// fetchExisting builds the path->blob reuse map and records every reachable +// subfs key so the old generation can be removed after publication. +func (c *client) fetchExisting(site string) (map[string]blobInfo, []string, error) { value, err := c.getRecord(c.did, collFs, site) if err != nil || value == nil { return map[string]blobInfo{}, nil, err } var rec FsRecord - if err := json.Unmarshal(value, &rec); err != nil || rec.Root == nil { + if err := json.Unmarshal(value, &rec); err != nil { + return nil, nil, fmt.Errorf("decode existing manifest: %w", err) + } + if rec.Root == nil { return map[string]blobInfo{}, nil, nil } - out := map[string]blobInfo{} - collectBlobs(rec.Root, "", out) - var uris []string - collectSubfsURIs(rec.Root, &uris) - for _, uri := range uris { - repo, rkey, ok := parseSubfsURI(uri) - if !ok { - continue - } - sv, err := c.getRecord(repo, collSubfs, rkey) - if err != nil || sv == nil { - continue + blobs := map[string]blobInfo{} + var rkeys []string + seen := map[string]bool{} + c.collectFsBlobs(rec.Root, "", blobs, &rkeys, seen) + sort.Strings(rkeys) + return blobs, rkeys, nil +} + +func (c *client) collectFsBlobs(dir *FsNode, prefix string, blobs map[string]blobInfo, rkeys *[]string, seen map[string]bool) { + for _, entry := range dir.Entries { + full := joinManifestPath(prefix, entry.Name) + switch entry.Node.NodeType { + case "file": + addBlob(blobs, full, entry.Node.Blob) + case "directory": + c.collectFsBlobs(entry.Node, full, blobs, rkeys, seen) + case "subfs": + subfsPrefix := full + if entry.Node.Flat != nil && *entry.Node.Flat { + subfsPrefix = prefix + } + c.collectSubfsBlobs(entry.Node.Subject, subfsPrefix, blobs, rkeys, seen) } - var sub SubfsRecord - if json.Unmarshal(sv, &sub) == nil && sub.Root != nil { - collectBlobs(sub.Root, "", out) + } +} + +func (c *client) collectSubfsBlobs(subject, prefix string, blobs map[string]blobInfo, rkeys *[]string, seen map[string]bool) { + repo, collection, rkey, ok := parseSubfsURI(subject) + if !ok || repo != c.did || collection != collSubfs || seen[subject] { + return + } + seen[subject] = true + *rkeys = append(*rkeys, rkey) + + value, err := c.getRecord(repo, collection, rkey) + if err != nil || value == nil { + return + } + var record SubfsRecord + if json.Unmarshal(value, &record) != nil || record.Root == nil { + return + } + c.walkSubfsNode(record.Root, prefix, blobs, rkeys, seen) +} + +func (c *client) walkSubfsNode(dir *SubfsNode, prefix string, blobs map[string]blobInfo, rkeys *[]string, seen map[string]bool) { + for _, entry := range dir.Entries { + full := joinManifestPath(prefix, entry.Name) + switch entry.Node.NodeType { + case "file": + addBlob(blobs, full, entry.Node.Blob) + case "directory": + c.walkSubfsNode(entry.Node, full, blobs, rkeys, seen) + case "subfs": + c.collectSubfsBlobs(entry.Node.Subject, prefix, blobs, rkeys, seen) } } - return out, rec.Root, nil } -func parseSubfsURI(uri string) (repo, rkey string, ok bool) { +func addBlob(blobs map[string]blobInfo, path string, raw json.RawMessage) { + if raw != nil { + blobs[path] = blobInfo{cid: blobRefLink(raw), blob: raw} + } +} + +func joinManifestPath(prefix, name string) string { + if prefix == "" { + return name + } + return prefix + "/" + name +} + +func parseSubfsURI(uri string) (repo, collection, rkey string, ok bool) { + if !strings.HasPrefix(uri, "at://") { + return "", "", "", false + } parts := strings.Split(strings.TrimPrefix(uri, "at://"), "/") - if len(parts) < 3 { - return "", "", false + if len(parts) != 3 || parts[0] == "" || parts[1] == "" || !validRecordKey(parts[2]) { + return "", "", "", false } - return parts[0], parts[2], true + return parts[0], parts[1], parts[2], true } type fileResult struct { @@ -132,95 +196,143 @@ type fileResult struct { cid string } -func (c *client) uploadBlobRetry(content []byte) (json.RawMessage, error) { +func (c *client) uploadBlobRetry(content []byte, onRateLimit func()) (json.RawMessage, error) { const retries = 3 var lastErr error - for attempt := 0; attempt < retries; attempt++ { + for attempt := range retries { blob, err := c.uploadBlob(content) if err == nil { return blob, nil } lastErr = err + if isStatus(err, http.StatusTooManyRequests) && onRateLimit != nil { + onRateLimit() + } if attempt == retries-1 { break } delay := time.Duration(500*(1< maxUploadConcurrency { + requested = maxUploadConcurrency + } + var largest int64 + for _, file := range files { + if file.size > largest { + largest = file.size + } + } + if largest > 0 { + byMemory := uploadMemoryBudget / (4 * largest) + if byMemory < 1 { + byMemory = 1 + } + if int64(requested) > byMemory { + requested = int(byMemory) + } + } + return requested +} + // processAndUpload compresses each file as needed, reuses existing blobs whose // CID is unchanged, and uploads the rest with bounded concurrency. func (c *client) processAndUpload(files []fileInfo, reuse map[string]blobInfo, concurrency int, useBase64, forceGzip bool, results map[string]*fileResult, mu *sync.Mutex) error { - if concurrency < 1 { - concurrency = 1 - } - sem := make(chan struct{}, concurrency) - var wg sync.WaitGroup - var errOnce sync.Once - var firstErr error - fail := func(e error) { errOnce.Do(func() { firstErr = e }) } - - for _, f := range files { - if firstErr != nil { - break + currentConcurrency := effectiveConcurrency(files, concurrency) + for start := 0; start < len(files); { + end := start + currentConcurrency + if end > len(files) { + end = len(files) } - wg.Add(1) - sem <- struct{}{} - go func(f fileInfo) { - defer wg.Done() - defer func() { <-sem }() - content, err := os.ReadFile(f.abs) - if err != nil { - fail(err) - return - } - mimeType := lookupMime(f.rel) - compress := forceGzip || shouldCompressFile(mimeType, f.rel) - - processed := content - base64Encoded := false - if compress { - gz, err := gzipBytes(content) - if err != nil { - fail(err) - return + var wg sync.WaitGroup + errs := make(chan error, end-start) + var rateLimited atomic.Bool + for _, file := range files[start:end] { + wg.Add(1) + go func(f fileInfo) { + defer wg.Done() + if err := c.processFile(f, reuse, useBase64, forceGzip, results, mu, func() { + rateLimited.Store(true) + }); err != nil { + errs <- err } - if !forceGzip && useBase64 && isTextMime(mimeType) { - processed = []byte(base64.StdEncoding.EncodeToString(gz)) - base64Encoded = true - } else { - processed = gz - } - } + }(file) + } + wg.Wait() + close(errs) + if err := <-errs; err != nil { + return err + } + if rateLimited.Load() && currentConcurrency > 2 { + currentConcurrency = 2 + fmt.Println("Rate limited; reducing upload concurrency to 2") + } + start = end + } + return nil +} - cid := computeCID(processed) - var blob json.RawMessage - if ex, ok := reuse[f.rel]; ok && ex.cid == cid { - blob = ex.blob - } else { - blob, err = c.uploadBlobRetry(processed) - if err != nil { - fail(fmt.Errorf("upload %s: %w", f.rel, err)) - return - } - } +func (c *client) processFile(f fileInfo, reuse map[string]blobInfo, useBase64, forceGzip bool, results map[string]*fileResult, mu *sync.Mutex, onRateLimit func()) error { + content, err := os.ReadFile(f.abs) + if err != nil { + return err + } + mimeType := lookupMime(f.rel) + compress := forceGzip || shouldCompressFile(mimeType, f.rel) - mu.Lock() - results[f.rel] = &fileResult{ - blob: blob, encoding: encodingFor(compress), mimeType: mimeType, - base64: base64Encoded, compress: compress, cid: cid, - } - mu.Unlock() - }(f) + processed := content + base64Encoded := false + if compress { + gz, err := gzipBytes(content) + if err != nil { + return err + } + if !forceGzip && useBase64 && isTextMime(mimeType) { + processed = []byte(base64.StdEncoding.EncodeToString(gz)) + base64Encoded = true + } else { + processed = gz + } + } + + cid := computeCID(processed) + var blob json.RawMessage + if existing, ok := reuse[f.rel]; ok && existing.cid == cid { + blob = existing.blob + } else { + blob, err = c.uploadBlobRetry(processed, onRateLimit) + if err != nil { + return fmt.Errorf("upload %s: %w", f.rel, err) + } + } + + mu.Lock() + results[f.rel] = &fileResult{ + blob: blob, encoding: encodingFor(compress), mimeType: mimeType, + base64: base64Encoded, compress: compress, cid: cid, } - wg.Wait() - return firstErr + mu.Unlock() + return nil } func encodingFor(compress bool) string { diff --git a/internal/wisp/upload_test.go b/internal/wisp/upload_test.go new file mode 100644 index 0000000..b8d830a --- /dev/null +++ b/internal/wisp/upload_test.go @@ -0,0 +1,156 @@ +package wisp + +import ( + "encoding/json" + "io" + "net/http" + "os" + "path/filepath" + "sync" + "testing" + "time" +) + +type rateLimitRT struct { + mu sync.Mutex + barrier chan struct{} + requests int + active int + maxLater int +} + +func newRateLimitRT() *rateLimitRT { + return &rateLimitRT{barrier: make(chan struct{})} +} + +func (rt *rateLimitRT) RoundTrip(req *http.Request) (*http.Response, error) { + body, err := io.ReadAll(req.Body) + if err != nil { + return nil, err + } + rt.mu.Lock() + rt.requests++ + requestNumber := rt.requests + rt.active++ + if requestNumber >= 5 && rt.active > rt.maxLater { + rt.maxLater = rt.active + } + if requestNumber == 3 { + close(rt.barrier) + } + rt.mu.Unlock() + + if requestNumber <= 3 { + <-rt.barrier + } + + rt.mu.Lock() + rt.active-- + rt.mu.Unlock() + if requestNumber == 1 { + return testResponse(http.StatusTooManyRequests, `{"error":"RateLimitExceeded"}`), nil + } + blob := `{"blob":{"ref":{"$link":"` + computeCID(body) + `"}}}` + return testResponse(http.StatusOK, blob), nil +} + +func TestProcessAndUploadReducesConcurrencyAfterRateLimit(t *testing.T) { + root := t.TempDir() + files := make([]fileInfo, 5) + for i := range files { + name := filepath.Join(root, string(rune('a'+i))+".bin") + if err := os.WriteFile(name, []byte{byte(i)}, 0o644); err != nil { + t.Fatal(err) + } + files[i] = fileInfo{abs: name, rel: filepath.Base(name), size: 1} + } + + rt := newRateLimitRT() + c := testClient(rt) + c.sleep = func(time.Duration) {} + results := map[string]*fileResult{} + var mu sync.Mutex + if err := c.processAndUpload(files, nil, 3, false, false, results, &mu); err != nil { + t.Fatal(err) + } + if len(results) != len(files) { + t.Fatalf("processed %d files, want %d", len(results), len(files)) + } + rt.mu.Lock() + defer rt.mu.Unlock() + if rt.requests != 6 { + t.Fatalf("requests=%d, want 6", rt.requests) + } + if rt.maxLater > 2 { + t.Fatalf("post-rate-limit concurrency=%d, want at most 2", rt.maxLater) + } +} + +func TestFetchExistingPreservesNestedSubfsPath(t *testing.T) { + const parentKey = "parent" + const chunkKey = "chunk" + blob := json.RawMessage(`{"ref":{"$link":"bafkreitest"}}`) + root := newFsDir() + root.Entries = append(root.Entries, &FsEntry{ + Name: "big", + Node: newFsSubfs("at://did:plc:test/place.wisp.subfs/"+parentKey, false), + }) + mainRecord := FsRecord{Type: collFs, Site: "site", Root: root, FileCount: 1} + parentRecord := SubfsRecord{ + Type: collSubfs, + Root: &SubfsNode{ + Type: "place.wisp.subfs#directory", + NodeType: "directory", + Entries: []*SubfsEntry{{ + Name: "chunk0", + Node: &SubfsNode{ + Type: "place.wisp.subfs#subfs", + NodeType: "subfs", + Subject: "at://did:plc:test/place.wisp.subfs/" + chunkKey, + }, + }}, + }, + } + chunkRecord := SubfsRecord{ + Type: collSubfs, + Root: &SubfsNode{ + Type: "place.wisp.subfs#directory", + NodeType: "directory", + Entries: []*SubfsEntry{{ + Name: "asset.txt", + Node: &SubfsNode{Type: "place.wisp.subfs#file", NodeType: "file", Blob: blob}, + }}, + }, + } + + rt := roundTripFunc(func(req *http.Request) (*http.Response, error) { + var record any + switch req.URL.Query().Get("rkey") { + case "site": + record = mainRecord + case parentKey: + record = parentRecord + case chunkKey: + record = chunkRecord + default: + t.Fatalf("unexpected rkey %q", req.URL.Query().Get("rkey")) + } + body, err := json.Marshal(map[string]any{"value": record}) + if err != nil { + t.Fatal(err) + } + return testResponse(http.StatusOK, string(body)), nil + }) + c := testClient(rt) + + blobs, rkeys, err := c.fetchExisting("site") + if err != nil { + t.Fatal(err) + } + if _, ok := blobs["big/asset.txt"]; !ok { + t.Fatalf("blob paths=%v", blobs) + } + if len(rkeys) != 2 || rkeys[0] != chunkKey || rkeys[1] != parentKey { + t.Fatalf("subfs keys=%v", rkeys) + } +} diff --git a/internal/wisp/util.go b/internal/wisp/util.go index 4bf1f62..cb67242 100644 --- a/internal/wisp/util.go +++ b/internal/wisp/util.go @@ -23,3 +23,19 @@ func formatBytes(n int64) string { } return fmt.Sprintf("%.1f %cB", float64(n)/float64(div), "KMGTPE"[exp]) } + +func validRecordKey(value string) bool { + if len(value) == 0 || len(value) > 512 || value == "." || value == ".." { + return false + } + for i := range len(value) { + ch := value[i] + if (ch < 'a' || ch > 'z') && + (ch < 'A' || ch > 'Z') && + (ch < '0' || ch > '9') && + ch != '.' && ch != '-' && ch != '_' && ch != ':' && ch != '~' { + return false + } + } + return true +}