feat(personal): add SSRF-safe web capture

This commit is contained in:
2026-07-12 04:26:40 +08:00
parent f2e87a2f35
commit a909806435
2 changed files with 755 additions and 0 deletions
@@ -0,0 +1,534 @@
package org.dromara.aihr.personal.service;
import org.dromara.aihr.personal.support.PersonalKnowledgeProperties;
import org.dromara.common.core.exception.ServiceException;
import org.springframework.stereotype.Service;
import javax.net.ssl.SNIHostName;
import javax.net.ssl.SSLParameters;
import javax.net.ssl.SSLSocket;
import javax.net.ssl.SSLSocketFactory;
import java.io.BufferedInputStream;
import java.io.ByteArrayOutputStream;
import java.io.IOException;
import java.io.InputStream;
import java.io.OutputStream;
import java.net.IDN;
import java.net.Inet4Address;
import java.net.Inet6Address;
import java.net.InetAddress;
import java.net.InetSocketAddress;
import java.net.Socket;
import java.net.URI;
import java.net.URISyntaxException;
import java.net.UnknownHostException;
import java.nio.charset.StandardCharsets;
import java.security.MessageDigest;
import java.security.NoSuchAlgorithmException;
import java.time.Instant;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.HashSet;
import java.util.HexFormat;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Locale;
import java.util.Map;
import java.util.Set;
@Service
public class PersonalUrlFetchService {
static final String BLOCKED = "PERSONAL_URL_BLOCKED";
static final String FETCH_FAILED = "PERSONAL_URL_FETCH_FAILED";
static final String RESPONSE_INVALID = "PERSONAL_URL_RESPONSE_INVALID";
static final String RESPONSE_TOO_LARGE = "PERSONAL_URL_RESPONSE_TOO_LARGE";
static final String CONTENT_TYPE_UNSUPPORTED = "PERSONAL_URL_CONTENT_TYPE_UNSUPPORTED";
static final String REDIRECT_LOOP = "PERSONAL_URL_REDIRECT_LOOP";
static final String REDIRECT_LIMIT = "PERSONAL_URL_REDIRECT_LIMIT";
private static final int MAX_URL_LENGTH = 4096;
private static final int MAX_REDIRECTS = 3;
private static final int MAX_HEADER_BYTES = 64 * 1024;
private static final int MAX_LINE_BYTES = 8 * 1024;
private static final long TOTAL_TIMEOUT_NANOS = 15_000_000_000L;
private static final String USER_AGENT = "wygj-personal-url-fetch/1.0";
private static final String ACCEPT = "text/html,text/plain,application/pdf,application/msword,application/vnd.ms-excel,application/vnd.ms-powerpoint,application/vnd.openxmlformats-officedocument.wordprocessingml.document,application/vnd.openxmlformats-officedocument.spreadsheetml.sheet,application/vnd.openxmlformats-officedocument.presentationml.presentation";
private static final Map<String, String> SAFE_HEADERS = Map.of(
"User-Agent", USER_AGENT,
"Accept", ACCEPT,
"Accept-Encoding", "identity"
);
private static final Set<String> ALLOWED_CONTENT_TYPES = Set.of(
"text/html", "text/plain", "text/markdown",
"application/pdf", "application/msword", "application/vnd.ms-excel",
"application/vnd.ms-powerpoint",
"application/vnd.openxmlformats-officedocument.wordprocessingml.document",
"application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
"application/vnd.openxmlformats-officedocument.presentationml.presentation"
);
private final PersonalKnowledgeProperties properties;
private final Resolver resolver;
private final Fetcher fetcher;
public PersonalUrlFetchService(PersonalKnowledgeProperties properties) {
this(properties, host -> Arrays.asList(InetAddress.getAllByName(host)), new RawSocketFetcher());
}
private PersonalUrlFetchService(PersonalKnowledgeProperties properties, Resolver resolver, Fetcher fetcher) {
this.properties = properties;
this.resolver = resolver;
this.fetcher = fetcher;
}
public static PersonalUrlFetchService forTest(PersonalKnowledgeProperties properties,
Resolver resolver, Fetcher fetcher) {
return new PersonalUrlFetchService(properties, resolver, fetcher);
}
/** Validate syntax, DNS answers and address policy. */
public URI validate(String rawUrl) {
return validateAndResolve(rawUrl).uri();
}
/** Capture a bounded public web resource without persisting it. */
public FetchResult fetch(String rawUrl) {
long maxBodyBytes = maxBodyBytes();
long deadline = System.nanoTime() + TOTAL_TIMEOUT_NANOS;
ValidatedTarget target = validateAndResolve(rawUrl);
Set<URI> visited = new HashSet<>();
visited.add(target.uri());
int redirects = 0;
while (true) {
TransportResponse response;
try {
response = fetcher.fetch(new FetchRequest(target.uri(), target.addresses(), deadline,
maxBodyBytes, SAFE_HEADERS));
} catch (ServiceException ex) {
throw ex;
} catch (Exception ex) {
throw new ServiceException(FETCH_FAILED);
}
if (response == null || response.body() == null || response.body().length > maxBodyBytes) {
throw new ServiceException(RESPONSE_TOO_LARGE);
}
enforceDeclaredLength(response.headers(), maxBodyBytes);
if (isRedirect(response.status())) {
if (redirects >= MAX_REDIRECTS) throw new ServiceException(REDIRECT_LIMIT);
String location = firstHeader(response.headers(), "location");
if (location == null || location.isBlank()) throw new ServiceException(RESPONSE_INVALID);
URI next;
try {
next = target.uri().resolve(location.trim());
} catch (IllegalArgumentException ex) {
throw new ServiceException(BLOCKED);
}
target = validateAndResolve(next.toString());
if (!visited.add(target.uri())) throw new ServiceException(REDIRECT_LOOP);
redirects++;
continue;
}
if (response.status() < 200 || response.status() >= 300) {
throw new ServiceException(RESPONSE_INVALID);
}
String contentType = normalizeContentType(firstHeader(response.headers(), "content-type"));
if (!ALLOWED_CONTENT_TYPES.contains(contentType)) {
throw new ServiceException(CONTENT_TYPE_UNSUPPORTED);
}
return new FetchResult(target.uri(), response.status(), contentType, response.body().clone(),
Instant.now(), sha256(response.body()));
}
}
private ValidatedTarget validateAndResolve(String rawUrl) {
URI uri = normalizeUri(rawUrl);
List<InetAddress> addresses;
try {
addresses = resolver.resolve(canonicalHost(uri));
} catch (Exception ex) {
throw new ServiceException(BLOCKED);
}
if (addresses == null || addresses.isEmpty()) throw new ServiceException(BLOCKED);
if (addresses.stream().anyMatch(address -> address == null || !isGloballyRoutable(address))) {
throw new ServiceException(BLOCKED);
}
List<InetAddress> copy = List.copyOf(addresses);
return new ValidatedTarget(uri, copy);
}
private static URI normalizeUri(String rawUrl) {
if (rawUrl == null || rawUrl.isBlank() || rawUrl.length() > MAX_URL_LENGTH) {
throw new ServiceException(BLOCKED);
}
try {
URI parsed = new URI(rawUrl.trim()).normalize();
String scheme = parsed.getScheme() == null ? "" : parsed.getScheme().toLowerCase(Locale.ROOT);
if (!("http".equals(scheme) || "https".equals(scheme))
|| parsed.getRawUserInfo() != null || parsed.getHost() == null || parsed.getHost().isBlank()) {
throw new ServiceException(BLOCKED);
}
int port = parsed.getPort();
if (port < -1 || port == 0 || port > 65535) throw new ServiceException(BLOCKED);
String rawHost = parsed.getHost();
if (rawHost.startsWith("[") && rawHost.endsWith("]")) rawHost = rawHost.substring(1, rawHost.length() - 1);
if (rawHost.indexOf('%') >= 0) throw new ServiceException(BLOCKED);
String host = rawHost.indexOf(':') >= 0 ? rawHost.toLowerCase(Locale.ROOT)
: IDN.toASCII(rawHost, IDN.USE_STD3_ASCII_RULES).toLowerCase(Locale.ROOT);
if (host.isBlank() || host.length() > 253) throw new ServiceException(BLOCKED);
if ("localhost".equals(host) || host.endsWith(".localhost")) throw new ServiceException(BLOCKED);
URI normalized = new URI(scheme, null, host, port,
parsed.getRawPath().isEmpty() ? "/" : parsed.getRawPath(), parsed.getRawQuery(), null).normalize();
if (normalized.toASCIIString().length() > MAX_URL_LENGTH) throw new ServiceException(BLOCKED);
return normalized;
} catch (URISyntaxException | IllegalArgumentException ex) {
if (ex instanceof ServiceException serviceException) throw serviceException;
throw new ServiceException(BLOCKED);
}
}
private long maxBodyBytes() {
try {
long value = Math.multiplyExact(properties.getMaxUrlBodyMb(), 1024L * 1024L);
if (value <= 0) throw new ArithmeticException();
return value;
} catch (ArithmeticException ex) {
throw new ServiceException(RESPONSE_TOO_LARGE);
}
}
private static boolean isRedirect(int status) {
return status == 301 || status == 302 || status == 303 || status == 307 || status == 308;
}
private static String canonicalHost(URI uri) {
String host = uri.getHost();
return host.startsWith("[") && host.endsWith("]") ? host.substring(1, host.length() - 1) : host;
}
private static String normalizeContentType(String value) {
if (value == null) return "";
int semicolon = value.indexOf(';');
return (semicolon < 0 ? value : value.substring(0, semicolon)).trim().toLowerCase(Locale.ROOT);
}
private static void enforceDeclaredLength(Map<String, List<String>> headers, long maxBodyBytes) {
String raw = firstHeader(headers, "content-length");
if (raw == null) return;
try {
long length = Long.parseLong(raw.trim());
if (length < 0 || length > maxBodyBytes) throw new ServiceException(RESPONSE_TOO_LARGE);
} catch (NumberFormatException ex) {
throw new ServiceException(RESPONSE_INVALID);
}
}
private static String firstHeader(Map<String, List<String>> headers, String name) {
if (headers == null) return null;
for (Map.Entry<String, List<String>> entry : headers.entrySet()) {
if (entry.getKey() != null && entry.getKey().equalsIgnoreCase(name)
&& entry.getValue() != null && !entry.getValue().isEmpty()) return entry.getValue().get(0);
}
return null;
}
static boolean isGloballyRoutable(InetAddress address) {
if (address.isAnyLocalAddress() || address.isLoopbackAddress() || address.isLinkLocalAddress()
|| address.isSiteLocalAddress() || address.isMulticastAddress()) return false;
byte[] bytes = address.getAddress();
if (address instanceof Inet4Address) return publicIpv4(bytes);
if (!(address instanceof Inet6Address) || bytes.length != 16) return false;
// Only global unicast 2000::/3, excluding IANA special-purpose prefixes below.
if ((bytes[0] & 0xe0) != 0x20) return false;
if (prefix(bytes, hex("20010000"), 23) || prefix(bytes, hex("20010db8"), 32)
|| prefix(bytes, hex("20020000"), 16) || prefix(bytes, hex("3fff0000"), 20)) return false;
return true;
}
private static boolean publicIpv4(byte[] bytes) {
if (bytes.length != 4) return false;
int a = bytes[0] & 255, b = bytes[1] & 255, c = bytes[2] & 255;
if (a == 0 || a == 10 || a == 127 || a >= 224) return false;
if (a == 100 && b >= 64 && b <= 127) return false;
if (a == 169 && b == 254) return false;
if (a == 172 && b >= 16 && b <= 31) return false;
if (a == 192 && (b == 168 || (b == 0 && (c == 0 || c == 2)))) return false;
if (a == 192 && ((b == 31 && c == 196) || (b == 52 && c == 193)
|| (b == 88 && c == 99) || (b == 175 && c == 48))) return false;
if (a == 198 && (b == 18 || b == 19 || (b == 51 && c == 100))) return false;
return !(a == 203 && b == 0 && c == 113);
}
private static boolean prefix(byte[] value, byte[] prefix, int bits) {
for (int i = 0; i < bits; i++) {
if (((value[i / 8] >> (7 - i % 8)) & 1) != ((prefix[i / 8] >> (7 - i % 8)) & 1)) return false;
}
return true;
}
private static byte[] hex(String value) {
return HexFormat.of().parseHex(value);
}
private static String sha256(byte[] body) {
try {
return HexFormat.of().formatHex(MessageDigest.getInstance("SHA-256").digest(body));
} catch (NoSuchAlgorithmException ex) {
throw new IllegalStateException("SHA-256 unavailable", ex);
}
}
@FunctionalInterface
public interface Resolver {
List<InetAddress> resolve(String host) throws UnknownHostException;
}
@FunctionalInterface
public interface Fetcher {
TransportResponse fetch(FetchRequest request) throws IOException;
}
public record FetchRequest(URI uri, List<InetAddress> addresses, long deadlineNanos,
long maxBodyBytes, Map<String, String> headers) {
public FetchRequest {
addresses = List.copyOf(addresses);
headers = Map.copyOf(headers);
}
}
public record TransportResponse(int status, Map<String, List<String>> headers, byte[] body) {
public TransportResponse {
headers = headers == null ? Map.of() : Map.copyOf(headers);
body = body == null ? new byte[0] : body.clone();
}
}
public record FetchResult(URI finalUri, int status, String contentType, byte[] body,
Instant capturedAt, String sha256) {
public FetchResult { body = body.clone(); }
@Override public byte[] body() { return body.clone(); }
}
private record ValidatedTarget(URI uri, List<InetAddress> addresses) { }
private static final class RawSocketFetcher implements Fetcher {
@Override
public TransportResponse fetch(FetchRequest request) throws IOException {
IOException last = null;
for (InetAddress address : request.addresses()) {
try {
return fetchAddress(request, address);
} catch (IOException ex) {
last = ex;
}
}
throw last == null ? new IOException("connection failed") : last;
}
private static TransportResponse fetchAddress(FetchRequest request, InetAddress address) throws IOException {
URI uri = request.uri();
int port = uri.getPort() >= 0 ? uri.getPort() : ("https".equals(uri.getScheme()) ? 443 : 80);
Socket plain = new Socket();
try {
plain.connect(new InetSocketAddress(address, port), timeout(request.deadlineNanos(), 5_000));
plain.setSoTimeout(timeout(request.deadlineNanos(), 5_000));
Socket active = plain;
if ("https".equals(uri.getScheme())) {
String tlsHost = canonicalHost(uri);
SSLSocket ssl = (SSLSocket) ((SSLSocketFactory) SSLSocketFactory.getDefault())
.createSocket(plain, tlsHost, port, true);
SSLParameters parameters = ssl.getSSLParameters();
parameters.setEndpointIdentificationAlgorithm("HTTPS");
if (!isIpLiteral(tlsHost)) parameters.setServerNames(List.of(new SNIHostName(tlsHost)));
ssl.setSSLParameters(parameters);
ssl.setSoTimeout(timeout(request.deadlineNanos(), 5_000));
ssl.startHandshake();
active = ssl;
}
writeRequest(active.getOutputStream(), request);
active.setSoTimeout(timeout(request.deadlineNanos(), 5_000));
TransportResponse response = parseHttpResponse(
new DeadlineInputStream(active.getInputStream(), active, request.deadlineNanos()), request.maxBodyBytes());
if (System.nanoTime() >= request.deadlineNanos()) throw new IOException("deadline exceeded");
return response;
} finally {
try { plain.close(); } catch (IOException ignored) { }
}
}
private static void writeRequest(OutputStream output, FetchRequest request) throws IOException {
URI uri = request.uri();
String target = uri.getRawPath().isEmpty() ? "/" : uri.getRawPath();
if (uri.getRawQuery() != null) target += "?" + uri.getRawQuery();
String host = hostHeader(uri);
StringBuilder value = new StringBuilder("GET ").append(target).append(" HTTP/1.1\r\nHost: ")
.append(host).append("\r\n");
request.headers().forEach((name, headerValue) -> value.append(name).append(": ")
.append(headerValue).append("\r\n"));
value.append("Connection: close\r\n\r\n");
output.write(value.toString().getBytes(StandardCharsets.US_ASCII));
output.flush();
}
private static String hostHeader(URI uri) {
String canonical = canonicalHost(uri);
String host = canonical.contains(":") ? "[" + canonical + "]" : canonical;
int defaultPort = "https".equals(uri.getScheme()) ? 443 : 80;
return uri.getPort() >= 0 && uri.getPort() != defaultPort ? host + ":" + uri.getPort() : host;
}
private static boolean isIpLiteral(String host) {
return host.indexOf(':') >= 0 || host.chars().allMatch(ch -> Character.isDigit(ch) || ch == '.');
}
private static int timeout(long deadlineNanos, int capMillis) throws IOException {
long remaining = deadlineNanos - System.nanoTime();
if (remaining <= 0) throw new IOException("deadline exceeded");
return (int) Math.max(1, Math.min(capMillis, (remaining + 999_999L) / 1_000_000L));
}
}
private static final class DeadlineInputStream extends InputStream {
private final InputStream delegate;
private final Socket socket;
private final long deadlineNanos;
private DeadlineInputStream(InputStream delegate, Socket socket, long deadlineNanos) {
this.delegate = delegate;
this.socket = socket;
this.deadlineNanos = deadlineNanos;
}
@Override
public int read() throws IOException {
socket.setSoTimeout(RawSocketFetcher.timeout(deadlineNanos, 5_000));
return delegate.read();
}
@Override
public int read(byte[] bytes, int offset, int length) throws IOException {
socket.setSoTimeout(RawSocketFetcher.timeout(deadlineNanos, 5_000));
return delegate.read(bytes, offset, length);
}
}
static TransportResponse parseHttpResponse(InputStream input, long maxBodyBytes) {
if (maxBodyBytes < 0) throw new ServiceException(RESPONSE_TOO_LARGE);
try {
BufferedInputStream buffered = input instanceof BufferedInputStream value ? value : new BufferedInputStream(input);
int[] headerBytes = {0};
String statusLine = readLine(buffered, headerBytes);
if (statusLine == null || !statusLine.startsWith("HTTP/1.")) throw new ServiceException(RESPONSE_INVALID);
String[] statusParts = statusLine.split(" ", 3);
if (statusParts.length < 2) throw new ServiceException(RESPONSE_INVALID);
int status;
try { status = Integer.parseInt(statusParts[1]); }
catch (NumberFormatException ex) { throw new ServiceException(RESPONSE_INVALID); }
Map<String, List<String>> headers = new LinkedHashMap<>();
while (true) {
String line = readLine(buffered, headerBytes);
if (line == null) throw new ServiceException(RESPONSE_INVALID);
if (line.isEmpty()) break;
int colon = line.indexOf(':');
if (colon <= 0) throw new ServiceException(RESPONSE_INVALID);
String name = line.substring(0, colon).trim().toLowerCase(Locale.ROOT);
String value = line.substring(colon + 1).trim();
if (name.isEmpty()) throw new ServiceException(RESPONSE_INVALID);
headers.computeIfAbsent(name, ignored -> new ArrayList<>()).add(value);
}
String transferEncoding = firstHeader(headers, "transfer-encoding");
String contentLength = firstHeader(headers, "content-length");
if (transferEncoding != null && contentLength != null) throw new ServiceException(RESPONSE_INVALID);
byte[] body;
if (transferEncoding != null) {
if (!"chunked".equalsIgnoreCase(transferEncoding.trim())) throw new ServiceException(RESPONSE_INVALID);
body = readChunked(buffered, maxBodyBytes);
} else if (contentLength != null) {
long length;
try { length = Long.parseLong(contentLength.trim()); }
catch (NumberFormatException ex) { throw new ServiceException(RESPONSE_INVALID); }
if (length < 0 || length > maxBodyBytes || length > Integer.MAX_VALUE) {
throw new ServiceException(RESPONSE_TOO_LARGE);
}
body = readExactly(buffered, (int) length);
} else {
body = readUntilEof(buffered, maxBodyBytes);
}
return new TransportResponse(status, headers, body);
} catch (ServiceException ex) {
throw ex;
} catch (IOException ex) {
throw new ServiceException(RESPONSE_INVALID);
}
}
private static byte[] readChunked(BufferedInputStream input, long maxBodyBytes) throws IOException {
ByteArrayOutputStream body = new ByteArrayOutputStream();
int[] framingBytes = {0};
while (true) {
String line = readLine(input, framingBytes);
if (line == null) throw new ServiceException(RESPONSE_INVALID);
int extension = line.indexOf(';');
String sizeText = (extension < 0 ? line : line.substring(0, extension)).trim();
long size;
try { size = Long.parseLong(sizeText, 16); }
catch (NumberFormatException ex) { throw new ServiceException(RESPONSE_INVALID); }
if (size < 0 || size > Integer.MAX_VALUE) throw new ServiceException(RESPONSE_INVALID);
if (size == 0) {
while (true) {
String trailer = readLine(input, framingBytes);
if (trailer == null) throw new ServiceException(RESPONSE_INVALID);
if (trailer.isEmpty()) return body.toByteArray();
if (trailer.indexOf(':') <= 0) throw new ServiceException(RESPONSE_INVALID);
}
}
if ((long) body.size() + size > maxBodyBytes) throw new ServiceException(RESPONSE_TOO_LARGE);
body.write(readExactly(input, (int) size));
if (input.read() != '\r' || input.read() != '\n') throw new ServiceException(RESPONSE_INVALID);
}
}
private static byte[] readExactly(InputStream input, int length) throws IOException {
byte[] bytes = input.readNBytes(length);
if (bytes.length != length) throw new ServiceException(RESPONSE_INVALID);
return bytes;
}
private static byte[] readUntilEof(InputStream input, long maxBodyBytes) throws IOException {
ByteArrayOutputStream body = new ByteArrayOutputStream();
byte[] buffer = new byte[8192];
int count;
while ((count = input.read(buffer)) >= 0) {
if ((long) body.size() + count > maxBodyBytes) throw new ServiceException(RESPONSE_TOO_LARGE);
body.write(buffer, 0, count);
}
return body.toByteArray();
}
private static String readLine(InputStream input, int[] totalBytes) throws IOException {
ByteArrayOutputStream line = new ByteArrayOutputStream();
int previous = -1;
while (true) {
int current = input.read();
if (current < 0) return line.size() == 0 && previous < 0 ? null : invalidLine();
totalBytes[0]++;
if (totalBytes[0] > MAX_HEADER_BYTES || line.size() > MAX_LINE_BYTES) {
throw new ServiceException(RESPONSE_INVALID);
}
if (previous == '\r') {
if (current != '\n') throw new ServiceException(RESPONSE_INVALID);
return line.toString(StandardCharsets.ISO_8859_1);
}
if (current == '\r') previous = current;
else {
if (current == '\n') throw new ServiceException(RESPONSE_INVALID);
line.write(current);
}
}
}
private static String invalidLine() {
throw new ServiceException(RESPONSE_INVALID);
}
}
@@ -0,0 +1,221 @@
package org.dromara.aihr.personal.service;
import org.dromara.aihr.personal.support.PersonalKnowledgeProperties;
import org.dromara.common.core.exception.ServiceException;
import org.junit.jupiter.api.Tag;
import org.junit.jupiter.api.Test;
import java.io.ByteArrayInputStream;
import java.net.InetAddress;
import java.net.URI;
import java.nio.charset.StandardCharsets;
import java.time.Instant;
import java.util.ArrayDeque;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.List;
import java.util.Map;
import static org.junit.jupiter.api.Assertions.*;
@Tag("dev")
class PersonalUrlFetchServiceTest {
private static final InetAddress PUBLIC = address("93.184.216.34");
@Test
void rejectsUnsafeSchemesSyntaxAndHosts() {
var service = fixture(host -> List.of(PUBLIC), request -> ok("text/plain", "ok"));
for (String raw : List.of(
"file:///etc/passwd", "ftp://example.com/a", "data:text/plain,hello",
"http://user:secret@example.com", "http:///missing", "not a url",
"http://localhost/admin", "http://service.localhost/admin",
"http://example.com/" + "x".repeat(5000))) {
assertCode("PERSONAL_URL_BLOCKED", () -> service.validate(raw), raw);
}
}
@Test
void rejectsUnsafeIpv4AndIpv6Ranges() {
for (String ip : List.of(
"0.0.0.1", "10.1.2.3", "100.64.0.1", "127.0.0.1", "169.254.169.254",
"172.16.0.1", "192.0.0.1", "192.0.2.1", "192.168.1.1", "198.18.0.1",
"198.51.100.1", "203.0.113.1", "224.0.0.1", "240.0.0.1", "255.255.255.255",
"::", "::1", "fe80::1", "fc00::1", "fd00::1", "ff02::1",
"2001:db8::1", "2001:2::1", "2001:100::1", "2002:5db8:d822::1",
"3fff::1", "64:ff9b::c0a8:101")) {
var service = fixture(host -> List.of(address(ip)), request -> ok("text/plain", "ok"));
assertCode("PERSONAL_URL_BLOCKED", () -> service.validate("http://example.com"), ip);
}
}
@Test
void rejectsEmptyOrMixedDnsAnswersAndAcceptsPublicResolution() {
assertCode("PERSONAL_URL_BLOCKED",
() -> fixture(host -> List.of(), request -> ok("text/plain", "ok")).validate("https://example.com"), "empty");
assertCode("PERSONAL_URL_BLOCKED",
() -> fixture(host -> List.of(PUBLIC, address("127.0.0.1")), request -> ok("text/plain", "ok"))
.validate("https://example.com"), "mixed");
assertCode("PERSONAL_URL_BLOCKED",
() -> fixture(host -> Arrays.asList(PUBLIC, null), request -> ok("text/plain", "ok"))
.validate("https://example.com"), "null answer");
assertEquals("https://example.com/a", fixture(host -> List.of(PUBLIC), request -> ok("text/plain", "ok"))
.validate("HTTPS://Example.COM/a").toString());
assertEquals("http://[2606:2800:220:1:248:1893:25c8:1946]/", fixture(
host -> List.of(address("2606:2800:220:1:248:1893:25c8:1946")), request -> ok("text/plain", "ok"))
.validate("http://[2606:2800:220:1:248:1893:25c8:1946]/").toString());
}
@Test
void followsRelativeRedirectAndRevalidatesEveryTarget() {
var seen = new ArrayList<PersonalUrlFetchService.FetchRequest>();
var responses = new ArrayDeque<PersonalUrlFetchService.TransportResponse>();
responses.add(response(302, Map.of("location", List.of("/final")), new byte[0]));
responses.add(ok("text/plain; charset=utf-8", "done"));
var service = fixture(host -> List.of(PUBLIC), request -> { seen.add(request); return responses.remove(); });
var result = service.fetch("https://example.com/start");
assertEquals(URI.create("https://example.com/final"), result.finalUri());
assertEquals("text/plain", result.contentType());
assertEquals("done", new String(result.body(), StandardCharsets.UTF_8));
assertEquals(2, seen.size());
assertEquals(List.of(PUBLIC), seen.get(0).addresses());
assertEquals(List.of(PUBLIC), seen.get(1).addresses());
}
@Test
void blocksUnsafeRedirectAndMixedAddressRedirect() {
var redirect = response(302, Map.of("location", List.of("http://metadata.test/latest")), new byte[0]);
var service = fixture(host -> host.equals("metadata.test")
? List.of(address("169.254.169.254")) : List.of(PUBLIC), request -> redirect);
assertCode("PERSONAL_URL_BLOCKED", () -> service.fetch("https://example.com/start"), "redirect private");
var mixed = fixture(host -> host.equals("mixed.test")
? List.of(PUBLIC, address("10.0.0.1")) : List.of(PUBLIC), request ->
response(301, Map.of("location", List.of("https://mixed.test/a")), new byte[0]));
assertCode("PERSONAL_URL_BLOCKED", () -> mixed.fetch("https://example.com"), "redirect mixed");
}
@Test
void detectsRedirectLoopAndMoreThanThreeRedirects() {
var loop = fixture(host -> List.of(PUBLIC), request ->
response(302, Map.of("location", List.of(request.uri().toString())), new byte[0]));
assertCode("PERSONAL_URL_REDIRECT_LOOP", () -> loop.fetch("https://example.com/a"), "loop");
var chain = fixture(host -> List.of(PUBLIC), request -> {
int n = Integer.parseInt(request.uri().getPath().substring(1));
return response(302, Map.of("location", List.of("/" + (n + 1))), new byte[0]);
});
assertCode("PERSONAL_URL_REDIRECT_LIMIT", () -> chain.fetch("https://example.com/0"), "limit");
}
@Test
void sendsOnlyFixedSafeHeaders() {
var requests = new ArrayList<PersonalUrlFetchService.FetchRequest>();
var service = fixture(host -> List.of(PUBLIC), request -> { requests.add(request); return ok("text/plain", "ok"); });
service.fetch("https://example.com/a");
Map<String, String> headers = requests.get(0).headers();
assertEquals(Map.of(
"User-Agent", "wygj-personal-url-fetch/1.0",
"Accept", "text/html,text/plain,application/pdf,application/msword,application/vnd.ms-excel,application/vnd.ms-powerpoint,application/vnd.openxmlformats-officedocument.wordprocessingml.document,application/vnd.openxmlformats-officedocument.spreadsheetml.sheet,application/vnd.openxmlformats-officedocument.presentationml.presentation",
"Accept-Encoding", "identity"), headers);
assertFalse(headers.keySet().stream().anyMatch(name -> List.of(
"cookie", "authorization", "proxy-authorization", "referer").contains(name.toLowerCase())));
}
@Test
void rejectsForbiddenOrMissingMimeAndOversizedBody() {
assertCode("PERSONAL_URL_CONTENT_TYPE_UNSUPPORTED",
() -> fixture(host -> List.of(PUBLIC), request -> ok("image/png", "x")).fetch("https://example.com"), "mime");
assertCode("PERSONAL_URL_CONTENT_TYPE_UNSUPPORTED",
() -> fixture(host -> List.of(PUBLIC), request -> response(200, Map.of(), "x".getBytes())).fetch("https://example.com"), "missing mime");
PersonalKnowledgeProperties properties = properties();
byte[] tooLarge = new byte[10 * 1024 * 1024 + 1];
assertCode("PERSONAL_URL_RESPONSE_TOO_LARGE", () ->
PersonalUrlFetchService.forTest(properties, host -> List.of(PUBLIC),
request -> ok("text/plain", tooLarge)).fetch("https://example.com"), "body cap");
}
@Test
void returnsDigestAndCaptureMetadata() {
var result = fixture(host -> List.of(PUBLIC), request -> ok("application/pdf", "abc"))
.fetch("https://example.com/a.pdf");
assertEquals(200, result.status());
assertEquals("application/pdf", result.contentType());
assertEquals("ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad", result.sha256());
assertTrue(result.capturedAt().isBefore(Instant.now().plusSeconds(1)));
}
@Test
void parsesBoundedContentLengthWithoutReadingOversizedBody() {
String raw = "HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: 11\r\n\r\nhello world";
assertEquals("hello world", new String(PersonalUrlFetchService.parseHttpResponse(
new ByteArrayInputStream(raw.getBytes(StandardCharsets.US_ASCII)), 11).body(), StandardCharsets.US_ASCII));
String oversized = "HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: 12\r\n\r\n";
assertCode("PERSONAL_URL_RESPONSE_TOO_LARGE", () -> PersonalUrlFetchService.parseHttpResponse(
new ByteArrayInputStream(oversized.getBytes(StandardCharsets.US_ASCII)), 11), "content length");
}
@Test
void parsesChunkedAndRejectsOverflowOrMalformedFraming() {
String valid = "HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nTransfer-Encoding: chunked\r\n\r\n5\r\nhello\r\n6\r\n world\r\n0\r\n\r\n";
assertEquals("hello world", new String(PersonalUrlFetchService.parseHttpResponse(
stream(valid), 11).body(), StandardCharsets.US_ASCII));
assertCode("PERSONAL_URL_RESPONSE_TOO_LARGE", () ->
PersonalUrlFetchService.parseHttpResponse(stream(valid), 10), "chunk overflow");
String malformed = "HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\nZZ\r\nhello\r\n0\r\n\r\n";
assertCode("PERSONAL_URL_RESPONSE_INVALID", () ->
PersonalUrlFetchService.parseHttpResponse(stream(malformed), 100), "chunk malformed");
}
@Test
void fetchRequestCarriesValidatedIpsAndOriginalTlsHost() {
var requests = new ArrayList<PersonalUrlFetchService.FetchRequest>();
fixture(host -> List.of(PUBLIC), request -> { requests.add(request); return ok("text/plain", "ok"); })
.fetch("https://example.com/path");
assertEquals("example.com", requests.get(0).uri().getHost());
assertEquals(List.of(PUBLIC), requests.get(0).addresses());
assertTrue(requests.get(0).deadlineNanos() > System.nanoTime());
}
private static PersonalUrlFetchService fixture(PersonalUrlFetchService.Resolver resolver,
PersonalUrlFetchService.Fetcher fetcher) {
return PersonalUrlFetchService.forTest(properties(), resolver, fetcher);
}
private static PersonalKnowledgeProperties properties() {
PersonalKnowledgeProperties properties = new PersonalKnowledgeProperties();
properties.setMaxUrlBodyMb(10);
return properties;
}
private static PersonalUrlFetchService.TransportResponse ok(String contentType, String body) {
return ok(contentType, body.getBytes(StandardCharsets.UTF_8));
}
private static PersonalUrlFetchService.TransportResponse ok(String contentType, byte[] body) {
return response(200, Map.of("content-type", List.of(contentType)), body);
}
private static PersonalUrlFetchService.TransportResponse response(int status, Map<String, List<String>> headers, byte[] body) {
return new PersonalUrlFetchService.TransportResponse(status, headers, body);
}
private static InetAddress address(String ip) {
try { return InetAddress.getByName(ip); }
catch (Exception ex) { throw new AssertionError(ex); }
}
private static ByteArrayInputStream stream(String value) {
return new ByteArrayInputStream(value.getBytes(StandardCharsets.US_ASCII));
}
private static void assertCode(String code, Runnable action, String context) {
assertEquals(code, assertThrows(ServiceException.class, action::run, context).getMessage(), context);
}
}