Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
63 changes: 50 additions & 13 deletions lib/r2/upload.test.ts
Original file line number Diff line number Diff line change
@@ -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(() => {
Expand All @@ -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",
);
Expand All @@ -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<string, string>;
};
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();
});
});
39 changes: 26 additions & 13 deletions lib/r2/upload.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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<string, string> = {
"image/jpeg": "jpg",
"image/png": "png",
Expand All @@ -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 {
Expand All @@ -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) },
Expand All @@ -88,9 +101,9 @@ export async function getPresignedUploadUrl(
* Expires in 1 hour.
*/
export async function getPresignedReadUrl(key: string): Promise<string> {
const client = getR2Client();
const { client, bucketName } = getR2Client();
const command = new GetObjectCommand({
Bucket: process.env.R2_BUCKET_NAME!,
Bucket: bucketName,
Key: key,
});

Expand Down