From a9098064355a258a0a77010939ac77cbc0e27681 Mon Sep 17 00:00:00 2001 From: let5sne Date: Sun, 12 Jul 2026 04:26:40 +0800 Subject: [PATCH] feat(personal): add SSRF-safe web capture --- .../service/PersonalUrlFetchService.java | 534 ++++++++++++++++++ .../personal/PersonalUrlFetchServiceTest.java | 221 ++++++++ 2 files changed, 755 insertions(+) create mode 100644 backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/service/PersonalUrlFetchService.java create mode 100644 backend/ruoyi-modules/ruoyi-aihr/src/test/java/org/dromara/aihr/personal/PersonalUrlFetchServiceTest.java diff --git a/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/service/PersonalUrlFetchService.java b/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/service/PersonalUrlFetchService.java new file mode 100644 index 00000000..834f398e --- /dev/null +++ b/backend/ruoyi-modules/ruoyi-aihr/src/main/java/org/dromara/aihr/personal/service/PersonalUrlFetchService.java @@ -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 SAFE_HEADERS = Map.of( + "User-Agent", USER_AGENT, + "Accept", ACCEPT, + "Accept-Encoding", "identity" + ); + private static final Set 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 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 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 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> 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> headers, String name) { + if (headers == null) return null; + for (Map.Entry> 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 resolve(String host) throws UnknownHostException; + } + + @FunctionalInterface + public interface Fetcher { + TransportResponse fetch(FetchRequest request) throws IOException; + } + + public record FetchRequest(URI uri, List addresses, long deadlineNanos, + long maxBodyBytes, Map headers) { + public FetchRequest { + addresses = List.copyOf(addresses); + headers = Map.copyOf(headers); + } + } + + public record TransportResponse(int status, Map> 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 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> 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); + } +} diff --git a/backend/ruoyi-modules/ruoyi-aihr/src/test/java/org/dromara/aihr/personal/PersonalUrlFetchServiceTest.java b/backend/ruoyi-modules/ruoyi-aihr/src/test/java/org/dromara/aihr/personal/PersonalUrlFetchServiceTest.java new file mode 100644 index 00000000..ff5f7a4a --- /dev/null +++ b/backend/ruoyi-modules/ruoyi-aihr/src/test/java/org/dromara/aihr/personal/PersonalUrlFetchServiceTest.java @@ -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(); + var responses = new ArrayDeque(); + 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(); + var service = fixture(host -> List.of(PUBLIC), request -> { requests.add(request); return ok("text/plain", "ok"); }); + service.fetch("https://example.com/a"); + + Map 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(); + 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> 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); + } +}