diff --git a/agent/src/main/java/dev/aikido/agent/Agent.java b/agent/src/main/java/dev/aikido/agent/Agent.java index d88d1df5..ccd68e4b 100644 --- a/agent/src/main/java/dev/aikido/agent/Agent.java +++ b/agent/src/main/java/dev/aikido/agent/Agent.java @@ -37,6 +37,7 @@ public static void premain(String agentArgs, Instrumentation inst) { } logger.info("Zen by Aikido v%s starting.", Config.pkgVersion); setAikidoSysProperties(); + PackageObserver.install(inst); // Test loading of zen binaries : loadLibrary(); @@ -59,6 +60,7 @@ public static void premain(String agentArgs, Instrumentation inst) { startDaemon(agentArgs); } + private static class AikidoTransformer { public static AgentBuilder.Transformer get() { var adviceAgentBuilder = new AgentBuilder.Transformer.ForAdvice() diff --git a/agent/src/main/java/dev/aikido/agent/PackageObserver.java b/agent/src/main/java/dev/aikido/agent/PackageObserver.java new file mode 100644 index 00000000..9f2bbf29 --- /dev/null +++ b/agent/src/main/java/dev/aikido/agent/PackageObserver.java @@ -0,0 +1,40 @@ +package dev.aikido.agent; + +import dev.aikido.agent_api.helpers.packages.RuntimePackageCollector; + +import java.lang.instrument.ClassFileTransformer; +import java.lang.instrument.Instrumentation; +import java.security.ProtectionDomain; + +final class PackageObserver implements ClassFileTransformer { + static void install(Instrumentation instrumentation) { + RuntimePackageCollector.start(); + instrumentation.addTransformer(new PackageObserver(), false); + for (Class loadedClass : instrumentation.getAllLoadedClasses()) { + observe(loadedClass); + } + } + + private static void observe(Class loadedClass) { + try { + RuntimePackageCollector.observeClass( + loadedClass.getName(), + loadedClass.getProtectionDomain() + ); + } catch (Throwable ignored) { + // Package reporting must never interfere with agent startup. + } + } + + @Override + public byte[] transform( + ClassLoader loader, + String className, + Class classBeingRedefined, + ProtectionDomain protectionDomain, + byte[] classfileBuffer + ) { + RuntimePackageCollector.observeClass(className, protectionDomain); + return null; + } +} diff --git a/agent_api/src/main/java/dev/aikido/agent_api/background/HeartbeatTask.java b/agent_api/src/main/java/dev/aikido/agent_api/background/HeartbeatTask.java index b9a95ab7..209a83b4 100644 --- a/agent_api/src/main/java/dev/aikido/agent_api/background/HeartbeatTask.java +++ b/agent_api/src/main/java/dev/aikido/agent_api/background/HeartbeatTask.java @@ -43,15 +43,17 @@ public void run() { Hostnames.HostnameEntry[] hostnames = HostnamesStore.getHostnamesAsList(); RouteEntry[] routes = RoutesStore.getRoutesAsList(); List users = UsersStore.getUsersAsList(); + List packages = RuntimePackagesStore.getPackagesAsList(); // Clear data : StatisticsStore.clear(); HostnamesStore.clear(); RoutesStore.clear(); UsersStore.clear(); + RuntimePackagesStore.clear(); // Create and send event : - Heartbeat.HeartbeatEvent event = Heartbeat.get(stats, hostnames, routes, users); + Heartbeat.HeartbeatEvent event = Heartbeat.get(stats, hostnames, routes, users, packages); Optional res = api.report(event); res.ifPresent(ServiceConfigStore::updateFromAPIResponse); } diff --git a/agent_api/src/main/java/dev/aikido/agent_api/background/cloud/api/events/Heartbeat.java b/agent_api/src/main/java/dev/aikido/agent_api/background/cloud/api/events/Heartbeat.java index d6082c4b..8ed3f405 100644 --- a/agent_api/src/main/java/dev/aikido/agent_api/background/cloud/api/events/Heartbeat.java +++ b/agent_api/src/main/java/dev/aikido/agent_api/background/cloud/api/events/Heartbeat.java @@ -2,6 +2,7 @@ import dev.aikido.agent_api.background.cloud.GetManagerInfo; import dev.aikido.agent_api.storage.Hostnames; +import dev.aikido.agent_api.storage.RuntimePackage; import dev.aikido.agent_api.storage.ServiceConfigStore; import dev.aikido.agent_api.storage.statistics.Statistics; import dev.aikido.agent_api.storage.routes.RouteEntry; @@ -22,15 +23,17 @@ public record HeartbeatEvent( Hostnames.HostnameEntry[] hostnames, RouteEntry[] routes, List users, + List packages, boolean middlewareInstalled ) implements APIEvent {} public static HeartbeatEvent get( - Statistics.StatsRecord stats, Hostnames.HostnameEntry[] hostnames, RouteEntry[] routes, List users + Statistics.StatsRecord stats, Hostnames.HostnameEntry[] hostnames, RouteEntry[] routes, + List users, List packages ) { long time = getUnixTimeMS(); // Get current time GetManagerInfo.ManagerInfo agent = getManagerInfo(); boolean middlewareInstalled = ServiceConfigStore.getConfig().isMiddlewareInstalled(); - return new HeartbeatEvent("heartbeat", agent, time, stats, hostnames, routes, users, middlewareInstalled); + return new HeartbeatEvent("heartbeat", agent, time, stats, hostnames, routes, users, packages, middlewareInstalled); } -} \ No newline at end of file +} diff --git a/agent_api/src/main/java/dev/aikido/agent_api/helpers/packages/JarPackageScanner.java b/agent_api/src/main/java/dev/aikido/agent_api/helpers/packages/JarPackageScanner.java new file mode 100644 index 00000000..d7cfbce9 --- /dev/null +++ b/agent_api/src/main/java/dev/aikido/agent_api/helpers/packages/JarPackageScanner.java @@ -0,0 +1,194 @@ +package dev.aikido.agent_api.helpers.packages; + +import dev.aikido.agent_api.storage.RuntimePackage; + +import java.io.BufferedInputStream; +import java.io.ByteArrayInputStream; +import java.io.IOException; +import java.io.InputStream; +import java.net.URI; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.Comparator; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Locale; +import java.util.Map; +import java.util.Properties; +import java.util.regex.Pattern; +import java.util.jar.JarEntry; +import java.util.jar.JarFile; +import java.util.jar.JarInputStream; + +public final class JarPackageScanner { + private static final int MAX_METADATA_BYTES = 1024 * 1024; + private static final Pattern MAVEN_PACKAGE_NAME = Pattern.compile("[A-Za-z0-9_.-]+:[A-Za-z0-9_.-]+"); + + private JarPackageScanner() {} + + public static List findMavenPackages( + String classResourceUrl, + long requiredAt + ) { + try { + JarLocation location = JarLocation.parse(classResourceUrl); + if (location == null) { + return List.of(); + } + if (location.nestedEntry() == null) { + try (InputStream input = new BufferedInputStream(Files.newInputStream(location.outerJar()))) { + return findMavenPackages(input, requiredAt); + } + } + try (JarFile outerJar = new JarFile(location.outerJar().toFile())) { + JarEntry nestedJar = outerJar.getJarEntry(location.nestedEntry()); + if (nestedJar == null) { + return List.of(); + } + try (InputStream input = new BufferedInputStream(outerJar.getInputStream(nestedJar))) { + return findMavenPackages(input, requiredAt); + } + } + } catch (IOException | RuntimeException ignored) { + return List.of(); + } + } + + public static String getJarLocationKey(String classResourceUrl) { + try { + JarLocation location = JarLocation.parse(classResourceUrl); + if (location == null) { + return null; + } + return location.getKey(); + } catch (RuntimeException ignored) { + return null; + } + } + + private static List findMavenPackages( + InputStream input, + long requiredAt + ) throws IOException { + Map packages = new LinkedHashMap<>(); + + try (JarInputStream jar = new JarInputStream(input)) { + JarEntry entry; + while ((entry = jar.getNextJarEntry()) != null) { + if (!entry.isDirectory() && isPomProperties(entry.getName())) { + byte[] metadata = jar.readNBytes(MAX_METADATA_BYTES + 1); + if (metadata.length <= MAX_METADATA_BYTES) { + addMavenPackage(metadata, requiredAt, packages); + } + } + } + } + + return packages.values().stream() + .sorted( + Comparator.comparing(RuntimePackage::name) + .thenComparing(RuntimePackage::version) + ) + .toList(); + } + + private static boolean isPomProperties(String name) { + return name.startsWith("META-INF/maven/") && name.endsWith("/pom.properties"); + } + + private static void addMavenPackage( + byte[] metadata, + long requiredAt, + Map packages + ) { + Properties properties = new Properties(); + try { + properties.load(new ByteArrayInputStream(metadata)); + } catch (IOException | IllegalArgumentException ignored) { + return; + } + String groupId = clean(properties.getProperty("groupId")); + String artifactId = clean(properties.getProperty("artifactId")); + String version = clean(properties.getProperty("version")); + if (groupId == null || artifactId == null || version == null) { + return; + } + String packageName = groupId + ":" + artifactId; + if (!MAVEN_PACKAGE_NAME.matcher(packageName).matches()) { + return; + } + add(packageName, version, requiredAt, packages); + } + + private static String clean(String value) { + if (value == null || value.isBlank()) { + return null; + } + return value.trim(); + } + + private static void add( + String name, + String version, + long requiredAt, + Map packages + ) { + packages.putIfAbsent(name + '\0' + version, new RuntimePackage(name, version, requiredAt)); + } + + private record JarLocation(Path outerJar, String nestedEntry) { + private static JarLocation parse(String url) { + if (url == null) { + return null; + } + String value = url; + if (value.startsWith("jar:")) { + value = value.substring(4); + } + if (value.startsWith("nested:")) { + value = value.substring(7); + } + String lowerCaseValue = value.toLowerCase(Locale.ROOT); + int outerEnd = lowerCaseValue.indexOf(".jar!/"); + int springBootOuterEnd = lowerCaseValue.indexOf(".jar/!"); + if (outerEnd < 0 || springBootOuterEnd >= 0 && springBootOuterEnd < outerEnd) { + outerEnd = springBootOuterEnd; + } + if (outerEnd < 0) { + int jarEnd = lowerCaseValue.indexOf(".jar"); + if (jarEnd < 0) { + return null; + } + Path jar = toPath(value.substring(0, jarEnd + 4)); + return new JarLocation(jar, null); + } + + Path outerJar = toPath(value.substring(0, outerEnd + 4)); + int nestedStart = outerEnd + 6; + int nestedEnd = lowerCaseValue.indexOf(".jar!/", nestedStart); + if (nestedEnd < 0 && lowerCaseValue.endsWith(".jar")) { + nestedEnd = value.length() - 4; + } + String nestedEntry = null; + if (nestedEnd >= 0) { + nestedEntry = value.substring(nestedStart, nestedEnd + 4); + } + return new JarLocation(outerJar, nestedEntry); + } + + private static Path toPath(String value) { + if (value.startsWith("file:")) { + return Path.of(URI.create(value)); + } + return Path.of(value); + } + + private String getKey() { + String key = outerJar.toAbsolutePath().normalize().toString(); + if (nestedEntry != null) { + key += "!/" + nestedEntry; + } + return key; + } + } +} diff --git a/agent_api/src/main/java/dev/aikido/agent_api/helpers/packages/RuntimePackageCollector.java b/agent_api/src/main/java/dev/aikido/agent_api/helpers/packages/RuntimePackageCollector.java new file mode 100644 index 00000000..6308d19c --- /dev/null +++ b/agent_api/src/main/java/dev/aikido/agent_api/helpers/packages/RuntimePackageCollector.java @@ -0,0 +1,71 @@ +package dev.aikido.agent_api.helpers.packages; + +import dev.aikido.agent_api.storage.RuntimePackagesStore; + +import java.security.ProtectionDomain; +import java.util.Set; +import java.util.concurrent.BlockingQueue; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.LinkedBlockingQueue; + +public final class RuntimePackageCollector { + private static final BlockingQueue PENDING_LOCATIONS = new LinkedBlockingQueue<>(); + private static final Set OBSERVED_LOCATIONS = ConcurrentHashMap.newKeySet(); + + private RuntimePackageCollector() {} + + public static void start() { + Thread worker = new Thread(RuntimePackageCollector::processLocations, "aikido-package-scanner"); + worker.setDaemon(true); + worker.start(); + } + + public static void observeClass(String className, ProtectionDomain protectionDomain) { + if (className == null || className.startsWith("dev/aikido/") || className.startsWith("dev.aikido.")) { + return; + } + try { + if (protectionDomain == null || protectionDomain.getCodeSource() == null) { + return; + } + String location = protectionDomain.getCodeSource().getLocation().toString(); + if (isAgentLocation(location)) { + return; + } + String locationKey = JarPackageScanner.getJarLocationKey(location); + if (locationKey != null && OBSERVED_LOCATIONS.add(locationKey)) { + PENDING_LOCATIONS.add(new ObservedLocation(location, System.currentTimeMillis())); + } + } catch (Throwable ignored) { + // Package reporting must never interfere with application class loading. + } + } + + private static void processLocations() { + while (!Thread.currentThread().isInterrupted()) { + try { + ObservedLocation location = PENDING_LOCATIONS.take(); + RuntimePackagesStore.addAll(JarPackageScanner.findMavenPackages(location.url(), location.requiredAt())); + } catch (InterruptedException interrupted) { + Thread.currentThread().interrupt(); + } catch (Throwable ignored) { + // A malformed or inaccessible JAR must not stop future package discovery. + } + } + } + + private static boolean isAgentLocation(String location) { + String agentDirectory = System.getProperty("AIK_agent_dir"); + if (agentDirectory == null) { + return false; + } + String directoryUrl = new java.io.File(agentDirectory).toURI().toString(); + return isAgentJar(location, directoryUrl); + } + + private static boolean isAgentJar(String url, String directoryUrl) { + return url.startsWith(directoryUrl + "agent.jar") || url.startsWith(directoryUrl + "agent_api.jar"); + } + + private record ObservedLocation(String url, long requiredAt) {} +} diff --git a/agent_api/src/main/java/dev/aikido/agent_api/storage/RuntimePackage.java b/agent_api/src/main/java/dev/aikido/agent_api/storage/RuntimePackage.java new file mode 100644 index 00000000..c9eb1635 --- /dev/null +++ b/agent_api/src/main/java/dev/aikido/agent_api/storage/RuntimePackage.java @@ -0,0 +1,3 @@ +package dev.aikido.agent_api.storage; + +public record RuntimePackage(String name, String version, long requiredAt) {} diff --git a/agent_api/src/main/java/dev/aikido/agent_api/storage/RuntimePackagesStore.java b/agent_api/src/main/java/dev/aikido/agent_api/storage/RuntimePackagesStore.java new file mode 100644 index 00000000..6d9e6be9 --- /dev/null +++ b/agent_api/src/main/java/dev/aikido/agent_api/storage/RuntimePackagesStore.java @@ -0,0 +1,36 @@ +package dev.aikido.agent_api.storage; + +import java.util.Comparator; +import java.util.Collection; +import java.util.List; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.ConcurrentMap; + +public final class RuntimePackagesStore { + private static final ConcurrentMap PACKAGES = new ConcurrentHashMap<>(); + + private RuntimePackagesStore() {} + + public static void addAll(Collection packages) { + for (RuntimePackage pkg : packages) { + PACKAGES.putIfAbsent(packageKey(pkg), pkg); + } + } + + public static List getPackagesAsList() { + return PACKAGES.values().stream() + .sorted( + Comparator.comparing(RuntimePackage::name) + .thenComparing(RuntimePackage::version) + ) + .toList(); + } + + public static void clear() { + PACKAGES.clear(); + } + + private static String packageKey(RuntimePackage pkg) { + return pkg.name() + '\0' + pkg.version(); + } +} diff --git a/agent_api/src/test/java/background/cloud/api/HeartbeatEventTest.java b/agent_api/src/test/java/background/cloud/api/HeartbeatEventTest.java index d2d92743..2179a419 100644 --- a/agent_api/src/test/java/background/cloud/api/HeartbeatEventTest.java +++ b/agent_api/src/test/java/background/cloud/api/HeartbeatEventTest.java @@ -55,7 +55,7 @@ public void testGetHeartbeatEvent() { ServiceConfigStore.setMiddlewareInstalled(false); // Act - Heartbeat.HeartbeatEvent event = Heartbeat.get(stats, hostnames.asArray(), routes, users); + Heartbeat.HeartbeatEvent event = Heartbeat.get(stats, hostnames.asArray(), routes, users, Collections.emptyList()); // Assert assertEquals("heartbeat", event.type()); @@ -68,7 +68,7 @@ public void testGetHeartbeatEvent() { // Test middleware installed as well : ServiceConfigStore.setMiddlewareInstalled(true); - Heartbeat.HeartbeatEvent event2 = Heartbeat.get(stats, hostnames.asArray(), routes, users); + Heartbeat.HeartbeatEvent event2 = Heartbeat.get(stats, hostnames.asArray(), routes, users, Collections.emptyList()); assertTrue(event2.middlewareInstalled()); } } diff --git a/agent_api/src/test/java/helpers/packages/JarPackageScannerTest.java b/agent_api/src/test/java/helpers/packages/JarPackageScannerTest.java new file mode 100644 index 00000000..f33501de --- /dev/null +++ b/agent_api/src/test/java/helpers/packages/JarPackageScannerTest.java @@ -0,0 +1,102 @@ +package helpers.packages; + +import dev.aikido.agent_api.helpers.packages.JarPackageScanner; +import dev.aikido.agent_api.storage.RuntimePackage; +import org.junit.jupiter.api.Test; + +import java.io.IOException; +import java.nio.charset.StandardCharsets; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.List; +import java.util.jar.Attributes; +import java.util.jar.JarEntry; +import java.util.jar.JarOutputStream; +import java.util.jar.Manifest; + +import static org.junit.jupiter.api.Assertions.assertEquals; + +class JarPackageScannerTest { + @Test + void readsMavenCoordinates() throws IOException { + Path jar = Files.createTempFile("aikido-package", ".jar"); + try (JarOutputStream output = new JarOutputStream(Files.newOutputStream(jar))) { + output.putNextEntry(new JarEntry("META-INF/maven/org.example/demo/pom.properties")); + output.write("groupId=org.example\nartifactId=demo\nversion=1.2.3\n" + .getBytes(StandardCharsets.UTF_8)); + } + + List packages = JarPackageScanner.findMavenPackages(jar.toUri().toString(), 123L); + + assertEquals(List.of(new RuntimePackage("org.example:demo", "1.2.3", 123L)), packages); + Files.deleteIfExists(jar); + } + + @Test + void readsMavenCoordinatesFromSpringBootNestedJar() throws IOException { + Path nestedJar = Files.createTempFile("aikido-nested-package", ".jar"); + try (JarOutputStream output = new JarOutputStream(Files.newOutputStream(nestedJar))) { + output.putNextEntry(new JarEntry("META-INF/maven/org.example/demo/pom.properties")); + output.write("groupId=org.example\nartifactId=demo\nversion=1.2.3\n" + .getBytes(StandardCharsets.UTF_8)); + } + + Path outerJar = Files.createTempFile("aikido-spring-boot", ".jar"); + try (JarOutputStream output = new JarOutputStream(Files.newOutputStream(outerJar))) { + output.putNextEntry(new JarEntry("BOOT-INF/lib/demo.jar")); + output.write(Files.readAllBytes(nestedJar)); + } + + String codeSourceUrl = "nested:" + outerJar + "/!BOOT-INF/lib/demo.jar"; + List packages = JarPackageScanner.findMavenPackages(codeSourceUrl, 123L); + + assertEquals(List.of(new RuntimePackage("org.example:demo", "1.2.3", 123L)), packages); + Files.deleteIfExists(nestedJar); + Files.deleteIfExists(outerJar); + } + + @Test + void ignoresManifestMetadataWithoutMavenCoordinates() throws IOException { + Path jar = Files.createTempFile("aikido-package", ".jar"); + Manifest manifest = new Manifest(); + manifest.getMainAttributes().put(Attributes.Name.MANIFEST_VERSION, "1.0"); + manifest.getMainAttributes().putValue("Implementation-Title", "demo"); + manifest.getMainAttributes().putValue("Implementation-Version", "2.0.0"); + try (JarOutputStream ignored = new JarOutputStream(Files.newOutputStream(jar), manifest)) {} + + List packages = JarPackageScanner.findMavenPackages(jar.toUri().toString(), 456L); + + assertEquals(List.of(), packages); + Files.deleteIfExists(jar); + } + + @Test + void ignoresIncompleteMavenCoordinates() throws IOException { + Path jar = Files.createTempFile("aikido-package", ".jar"); + try (JarOutputStream output = new JarOutputStream(Files.newOutputStream(jar))) { + output.putNextEntry(new JarEntry("META-INF/maven/unknown/demo/pom.properties")); + output.write("artifactId=demo\nversion=1.2.3\n".getBytes(StandardCharsets.UTF_8)); + } + + List packages = JarPackageScanner.findMavenPackages(jar.toUri().toString(), 123L); + + assertEquals(List.of(), packages); + Files.deleteIfExists(jar); + } + + @Test + void ignoresUnresolvedMavenProperties() throws IOException { + Path jar = Files.createTempFile("aikido-package", ".jar"); + try (JarOutputStream output = new JarOutputStream(Files.newOutputStream(jar))) { + output.putNextEntry(new JarEntry("META-INF/maven/unknown/demo/pom.properties")); + output.write(("groupId=${project.groupId}\n" + + "artifactId=demo\n" + + "version=1.2.3\n").getBytes(StandardCharsets.UTF_8)); + } + + List packages = JarPackageScanner.findMavenPackages(jar.toUri().toString(), 123L); + + assertEquals(List.of(), packages); + Files.deleteIfExists(jar); + } +} diff --git a/end2end/spring_boot_postgres.py b/end2end/spring_boot_postgres.py index 8421b6c1..b182ee5f 100644 --- a/end2end/spring_boot_postgres.py +++ b/end2end/spring_boot_postgres.py @@ -19,4 +19,5 @@ unsafe_request=Request("/api/files/read", data_type='form', body={"fileName": "./../databases/docker-compose.yml"}) ) +spring_boot_postgres_app.test_dependency_detection("com.fasterxml.jackson.core:jackson-databind") spring_boot_postgres_app.test_all_payloads() diff --git a/end2end/utils/EventHandler.py b/end2end/utils/EventHandler.py index 5edfb7fa..918b899b 100644 --- a/end2end/utils/EventHandler.py +++ b/end2end/utils/EventHandler.py @@ -15,6 +15,8 @@ def fetch_events_from_mock(self): return json_events def fetch_attacks(self): return filter_on_event_type(self.fetch_events_from_mock(), "detected_attack") + def fetch_heartbeats(self): + return filter_on_event_type(self.fetch_events_from_mock(), "heartbeat") def set_protection(self, api_pets_create_protection, api_protection): print("Setting forceProtectionOff") res = requests.post(self.url + "/mock/set_protection", json={ @@ -23,4 +25,4 @@ def set_protection(self, api_pets_create_protection, api_protection): }, timeout=5) def filter_on_event_type(events, type): - return [event for event in events if event["type"] == type] \ No newline at end of file + return [event for event in events if event["type"] == type] diff --git a/end2end/utils/__init__.py b/end2end/utils/__init__.py index 4afd6342..2a1fddb4 100644 --- a/end2end/utils/__init__.py +++ b/end2end/utils/__init__.py @@ -62,6 +62,17 @@ def test_all_payloads(self): for key in self.payloads.keys(): self.test_payload(key) + def test_dependency_detection(self, dependency_name, timeout=75): + deadline = time.time() + timeout + while time.time() < deadline: + for heartbeat in self.event_handler.fetch_heartbeats(): + packages = heartbeat.get("packages", []) + if any(package.get("name") == dependency_name for package in packages): + print("✅ Detected loaded dependency: " + dependency_name) + return + time.sleep(1) + raise AssertionError("Loaded dependency was not reported: " + dependency_name) + def test_blocking(self): test_bot_blocking(self.urls["enabled"]) print("✅ Tested bot blocking")