diff --git a/lib/r2/upload.test.ts b/lib/r2/upload.test.ts index fe79b6c..5e672a3 100644 --- a/lib/r2/upload.test.ts +++ b/lib/r2/upload.test.ts @@ -1,32 +1,39 @@ import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; -const { putObjectCommand, getSignedUrl } = vi.hoisted(() => ({ - putObjectCommand: vi.fn((input: unknown) => ({ input })), - getSignedUrl: vi.fn(async () => "https://upload.example/signed"), -})); +const { s3Client, putObjectCommand, getObjectCommand, getSignedUrl } = + vi.hoisted(() => ({ + s3Client: vi.fn(() => ({})), + putObjectCommand: vi.fn((input: unknown) => ({ input })), + getObjectCommand: vi.fn((input: unknown) => ({ input })), + getSignedUrl: vi.fn(async () => "https://upload.example/signed"), + })); vi.mock("@aws-sdk/client-s3", () => ({ - S3Client: vi.fn(() => ({})), + S3Client: s3Client, PutObjectCommand: putObjectCommand, - GetObjectCommand: vi.fn((input: unknown) => ({ input })), + GetObjectCommand: getObjectCommand, })); vi.mock("@aws-sdk/s3-request-presigner", () => ({ getSignedUrl, })); -import { getPresignedUploadUrl } from "./upload"; - describe("getPresignedUploadUrl", () => { const originalEnv = { ...process.env }; + let upload: typeof import("./upload"); - beforeEach(() => { + beforeEach(async () => { + vi.resetModules(); process.env.R2_ACCOUNT_ID = "test-account"; process.env.R2_ACCESS_KEY_ID = "test-key"; process.env.R2_SECRET_ACCESS_KEY = "test-secret"; process.env.R2_BUCKET_NAME = "test-bucket"; process.env.R2_PUBLIC_URL = "https://cdn.example.com"; + s3Client.mockClear(); putObjectCommand.mockClear(); + getObjectCommand.mockClear(); + getSignedUrl.mockClear(); + upload = await import("./upload"); }); afterEach(() => { @@ -39,20 +46,26 @@ describe("getPresignedUploadUrl", () => { ["image/webp", "webp"], ["image/gif", "gif"], ])("maps %s to a .%s key", async (contentType, ext) => { - const { key } = await getPresignedUploadUrl("whatever.html", contentType); + const { key } = await upload.getPresignedUploadUrl( + "whatever.html", + contentType, + ); expect(key).toMatch(new RegExp(`^screenshots/[^/]+\\.${ext}$`)); }); it("ignores a hostile filename with a disallowed extension", async () => { - const { key } = await getPresignedUploadUrl("payload.html", "image/png"); + const { key } = await upload.getPresignedUploadUrl( + "payload.html", + "image/png", + ); expect(key.endsWith(".html")).toBe(false); expect(key.endsWith(".png")).toBe(true); }); it("falls back to .bin for an unrecognized content type", async () => { - const { key } = await getPresignedUploadUrl( + const { key } = await upload.getPresignedUploadUrl( "file", "application/octet-stream", ); @@ -61,11 +74,35 @@ describe("getPresignedUploadUrl", () => { }); it("passes the sanitized original filename as object metadata, not the key", async () => { - await getPresignedUploadUrl("my photo #1!.png", "image/png"); + await upload.getPresignedUploadUrl("my photo #1!.png", "image/png"); const command = putObjectCommand.mock.calls[0][0] as { Metadata?: Record; }; expect(command.Metadata?.["original-filename"]).toBe("my_photo__1_.png"); }); + + it("reuses one client for upload and read presigns", async () => { + await upload.getPresignedUploadUrl("first.png", "image/png"); + await upload.getPresignedReadUrl("screenshots/existing.png"); + await upload.getPresignedUploadUrl("second.webp", "image/webp"); + + expect(s3Client).toHaveBeenCalledTimes(1); + expect(getSignedUrl).toHaveBeenCalledTimes(3); + }); + + it.each([ + "R2_ACCOUNT_ID", + "R2_ACCESS_KEY_ID", + "R2_SECRET_ACCESS_KEY", + "R2_BUCKET_NAME", + ])("names a missing %s variable before signing", async (variable) => { + delete process.env[variable]; + + await expect( + upload.getPresignedReadUrl("screenshots/existing.png"), + ).rejects.toThrow(`Missing ${variable} environment variable.`); + expect(s3Client).not.toHaveBeenCalled(); + expect(getSignedUrl).not.toHaveBeenCalled(); + }); }); diff --git a/lib/r2/upload.ts b/lib/r2/upload.ts index 04b539f..1d64c09 100644 --- a/lib/r2/upload.ts +++ b/lib/r2/upload.ts @@ -7,8 +7,6 @@ import { getSignedUrl } from "@aws-sdk/s3-request-presigner"; import { randomUUID } from "crypto"; import { getR2PublicBaseUrl } from "@/lib/utils/urls"; -const R2_ENDPOINT = `https://${process.env.R2_ACCOUNT_ID}.r2.cloudflarestorage.com`; - const EXT_BY_CONTENT_TYPE: Record = { "image/jpeg": "jpg", "image/png": "png", @@ -30,15 +28,30 @@ function sanitizeFilenameForMetadata(filename: string): string { return filename.replace(/[^\w.-]/g, "_").slice(0, 255); } -function getR2Client() { - return new S3Client({ +let r2: { client: S3Client; bucketName: string } | undefined; + +function requiredEnv(name: string): string { + const value = process.env[name]; + if (!value?.trim()) { + throw new Error(`Missing ${name} environment variable.`); + } + return value; +} + +function getR2Client(): { client: S3Client; bucketName: string } { + if (r2) return r2; + + const accountId = requiredEnv("R2_ACCOUNT_ID"); + const accessKeyId = requiredEnv("R2_ACCESS_KEY_ID"); + const secretAccessKey = requiredEnv("R2_SECRET_ACCESS_KEY"); + const bucketName = requiredEnv("R2_BUCKET_NAME"); + const client = new S3Client({ region: "auto", - endpoint: R2_ENDPOINT, - credentials: { - accessKeyId: process.env.R2_ACCESS_KEY_ID!, - secretAccessKey: process.env.R2_SECRET_ACCESS_KEY!, - }, + endpoint: `https://${accountId}.r2.cloudflarestorage.com`, + credentials: { accessKeyId, secretAccessKey }, }); + r2 = { client, bucketName }; + return r2; } interface PresignUploadResult { @@ -62,9 +75,9 @@ export async function getPresignedUploadUrl( const ext = extensionForContentType(contentType); const key = `${folder}/${randomUUID()}.${ext}`; - const client = getR2Client(); + const { client, bucketName } = getR2Client(); const command = new PutObjectCommand({ - Bucket: process.env.R2_BUCKET_NAME!, + Bucket: bucketName, Key: key, ContentType: contentType, Metadata: { "original-filename": sanitizeFilenameForMetadata(filename) }, @@ -88,9 +101,9 @@ export async function getPresignedUploadUrl( * Expires in 1 hour. */ export async function getPresignedReadUrl(key: string): Promise { - const client = getR2Client(); + const { client, bucketName } = getR2Client(); const command = new GetObjectCommand({ - Bucket: process.env.R2_BUCKET_NAME!, + Bucket: bucketName, Key: key, });