diff --git a/src/main/java/com/volcengine/ark/runtime/selfhosted/EnvironmentWorker.java b/src/main/java/com/volcengine/ark/runtime/selfhosted/EnvironmentWorker.java index 381d89c..9525236 100644 --- a/src/main/java/com/volcengine/ark/runtime/selfhosted/EnvironmentWorker.java +++ b/src/main/java/com/volcengine/ark/runtime/selfhosted/EnvironmentWorker.java @@ -9,11 +9,9 @@ import java.io.IOException; import java.lang.management.ManagementFactory; import java.net.InetAddress; -import java.nio.charset.StandardCharsets; import java.nio.file.Files; import java.nio.file.Path; import java.nio.file.Paths; -import java.security.MessageDigest; import java.util.LinkedHashMap; import java.util.Map; import java.util.concurrent.atomic.AtomicBoolean; @@ -61,7 +59,7 @@ public void run() { return; } try { - handleItem(claimedWorkFromItem(item), false); + handleItem(claimedWorkFromItem(item)); } catch (SessionToolRunner.IdleTimeoutException | SessionToolRunner.SessionTerminatedException ignored) { } catch (Exception e) { options.logger.log(Level.WARNING, "handle work failed", e); @@ -77,22 +75,23 @@ public void handleItem(HandleItemOptions handleOptions) throws IOException { Thread previous = activeThread; activeThread = Thread.currentThread(); try { - handleItem(claimedWorkFromOptions(handleOptions), true); + handleItem(claimedWorkFromOptions(handleOptions)); } catch (SessionToolRunner.IdleTimeoutException | SessionToolRunner.SessionTerminatedException ignored) { } finally { activeThread = previous; } } - private void handleItem(ClaimedWork work, boolean useWorkdirAsSession) throws IOException { + private void handleItem(ClaimedWork work) throws IOException { if (work.environmentId.isEmpty()) { work.environmentId = firstNonEmpty(options.environmentId, System.getenv("MA_ENVIRONMENT_ID")); } AtomicBoolean stop = new AtomicBoolean(false); AtomicReference heartbeatCause = new AtomicReference<>(""); Thread heartbeat = null; + Initializer initializer = null; try { - String workdir = workdirFor(work.sessionId, useWorkdirAsSession); + String workdir = workdir(); Thread heartbeatThread = new Thread( () -> heartbeatLoop(work, stop, heartbeatCause), "ma-self-host-heartbeat"); heartbeatThread.setDaemon(true); @@ -108,12 +107,13 @@ private void handleItem(ClaimedWork work, boolean useWorkdirAsSession) throws IO if (session.getId() == null || session.getId().isEmpty()) { session.setId(work.sessionId); } - new Initializer(api, new Initializer.Options(workdir)).setup(session); + initializer = new Initializer(api, new Initializer.Options(workdir)); + initializer.setup(session); if (closed.get() || stop.get()) { return; } ToolContext toolContext = toolContext(workdir, stop); - FileToolResultStore store = new FileToolResultStore(workdir); + FileToolResultStore store = new FileToolResultStore(workdir, work.sessionId); SessionToolRunner runner = new SessionToolRunner(api, work.sessionId, new SessionToolRunner.Options() .workId(work.id) .tools(options.tools == null ? DefaultTools.create() : options.tools) @@ -130,6 +130,13 @@ private void handleItem(ClaimedWork work, boolean useWorkdirAsSession) throws IO activeRunner = null; } } finally { + if (initializer != null) { + try { + initializer.cleanup(); + } catch (IOException error) { + options.logger.log(Level.WARNING, "cleanup session skills failed", error); + } + } stop.set(true); if (heartbeat != null) { try { @@ -243,17 +250,12 @@ private ToolContext toolContext(String workdir, AtomicBoolean workStop) { return copy; } - private String workdirFor(String sessionId, boolean useWorkdirAsSession) throws IOException { + private String workdir() throws IOException { Path root = Paths.get(options.workdir == null || options.workdir.isEmpty() ? "." : options.workdir) .toAbsolutePath() .normalize(); Files.createDirectories(root); - if (useWorkdirAsSession) { - return root.toString(); - } - Path sessionDir = root.resolve(sessionWorkdirName(sessionId)).normalize(); - Files.createDirectories(sessionDir); - return sessionDir.toString(); + return root.toString(); } private ClaimedWork claimedWorkFromOptions(HandleItemOptions opts) { @@ -339,26 +341,6 @@ private static boolean shouldStopItem(String heartbeatCause) { && !"heartbeat_permanent_failure".equals(heartbeatCause); } - private static String sessionWorkdirName(String sessionId) { - if (sessionId != null - && sessionId.matches("[A-Za-z0-9._-]+") - && !".".equals(sessionId) - && !"..".equals(sessionId)) { - return sessionId; - } - try { - byte[] digest = MessageDigest.getInstance("SHA-256") - .digest(String.valueOf(sessionId).getBytes(StandardCharsets.UTF_8)); - StringBuilder value = new StringBuilder("session-"); - for (byte item : digest) { - value.append(String.format("%02x", item)); - } - return value.toString(); - } catch (Exception error) { - throw new IllegalStateException("failed to hash session id", error); - } - } - private static String firstNonEmpty(String first, String second) { return first != null && !first.isEmpty() ? first : (second == null ? "" : second); } diff --git a/src/main/java/com/volcengine/ark/runtime/selfhosted/FileToolResultStore.java b/src/main/java/com/volcengine/ark/runtime/selfhosted/FileToolResultStore.java index cfece6f..a2e1df9 100644 --- a/src/main/java/com/volcengine/ark/runtime/selfhosted/FileToolResultStore.java +++ b/src/main/java/com/volcengine/ark/runtime/selfhosted/FileToolResultStore.java @@ -32,10 +32,21 @@ public class FileToolResultStore { private final Path dir; public FileToolResultStore(String workdir) throws IOException { + this(workdir, null); + } + + public FileToolResultStore(String workdir, String sessionId) throws IOException { if (workdir == null || workdir.isEmpty()) { throw new IllegalArgumentException("workdir must not be empty"); } - this.dir = Paths.get(workdir, ".ma_self_host_worker", "tool_ledger"); + Path storeDir = Paths.get(workdir, ".ma_self_hosted_worker", "tool_ledger"); + if (sessionId != null) { + if (sessionId.isEmpty()) { + throw new IllegalArgumentException("session id must not be empty"); + } + storeDir = storeDir.resolve(sessionLedgerName(sessionId)); + } + this.dir = storeDir; Files.createDirectories(this.dir); } @@ -196,6 +207,13 @@ private static String sha256(String value) { } } + private static String sessionLedgerName(String sessionId) { + if (sessionId.matches("[A-Za-z0-9._-]+") && !".".equals(sessionId) && !"..".equals(sessionId)) { + return sessionId; + } + return "session-" + sha256(sessionId); + } + @SuppressWarnings("unchecked") private static Map asMap(Object value) { return value instanceof Map ? (Map) value : Collections.emptyMap(); diff --git a/src/main/java/com/volcengine/ark/runtime/selfhosted/Initializer.java b/src/main/java/com/volcengine/ark/runtime/selfhosted/Initializer.java index 0e3101e..50d9d3d 100644 --- a/src/main/java/com/volcengine/ark/runtime/selfhosted/Initializer.java +++ b/src/main/java/com/volcengine/ark/runtime/selfhosted/Initializer.java @@ -11,6 +11,7 @@ import java.nio.file.Paths; import java.nio.file.SimpleFileVisitor; import java.nio.file.attribute.BasicFileAttributes; +import java.util.ArrayList; import java.util.List; import java.util.logging.Level; import java.util.logging.Logger; @@ -26,6 +27,7 @@ public class Initializer { private final SelfHostedClient api; private final Options options; + private final List installedSkillDirs = new ArrayList<>(); public Initializer(SelfHostedClient api, Options options) { if (api == null) { @@ -70,6 +72,25 @@ public void setup(SessionSnapshot session) throws IOException { } } + public void cleanup() throws IOException { + IOException firstError = null; + for (Path path : installedSkillDirs) { + try { + deleteRecursively(path); + } catch (IOException error) { + if (firstError == null) { + firstError = error; + } else { + firstError.addSuppressed(error); + } + } + } + installedSkillDirs.clear(); + if (firstError != null) { + throw firstError; + } + } + public void installSkill(String sessionId, SkillRef skill) throws IOException { Files.createDirectories(Paths.get(options.workdir)); Files.createDirectories(Paths.get(options.skillsDir)); @@ -91,6 +112,7 @@ public void installSkill(String sessionId, SkillRef skill) throws IOException { Path source = installSourceDir(tmp); Path target = Paths.get(options.skillsDir, name); Path backup = replaceSkillDir(source, target); + installedSkillDirs.add(target); committed = true; if (!source.equals(tmp)) { try { diff --git a/src/test/java/com/volcengine/ark/runtime/selfhosted/EnvironmentWorkerTest.java b/src/test/java/com/volcengine/ark/runtime/selfhosted/EnvironmentWorkerTest.java index 71715c2..c654d49 100644 --- a/src/test/java/com/volcengine/ark/runtime/selfhosted/EnvironmentWorkerTest.java +++ b/src/test/java/com/volcengine/ark/runtime/selfhosted/EnvironmentWorkerTest.java @@ -11,6 +11,7 @@ import java.io.IOException; import java.lang.reflect.Method; import java.nio.file.Files; +import java.nio.file.Path; import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; @@ -24,6 +25,18 @@ import org.junit.Test; public class EnvironmentWorkerTest { + @Test + public void workerUsesConfiguredWorkdir() throws Exception { + Path workdir = Files.createTempDirectory("ark-java-worker-"); + EnvironmentWorker worker = new EnvironmentWorker( + new SelfHostedClient("test-key"), + new EnvironmentWorker.Options().workdir(workdir.toString())); + Method method = EnvironmentWorker.class.getDeclaredMethod("workdir"); + method.setAccessible(true); + + assertEquals(workdir.toAbsolutePath().normalize().toString(), method.invoke(worker)); + } + @Test public void emptyHeartbeatResponseDoesNotSpin() throws Exception { NullHeartbeatClient client = new NullHeartbeatClient(); diff --git a/src/test/java/com/volcengine/ark/runtime/selfhosted/FileToolResultStoreTest.java b/src/test/java/com/volcengine/ark/runtime/selfhosted/FileToolResultStoreTest.java index a2b36c1..0da1be3 100644 --- a/src/test/java/com/volcengine/ark/runtime/selfhosted/FileToolResultStoreTest.java +++ b/src/test/java/com/volcengine/ark/runtime/selfhosted/FileToolResultStoreTest.java @@ -5,6 +5,7 @@ import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertFalse; +import static org.junit.Assert.assertTrue; import java.nio.file.Files; import java.nio.file.Path; @@ -32,7 +33,7 @@ public void recoveryRemovesStaleTemporaryRecords() throws Exception { Path workdir = Files.createTempDirectory("ark-java-store-"); FileToolResultStore store = new FileToolResultStore(workdir.toString()); Path stale = workdir - .resolve(".ma_self_host_worker") + .resolve(".ma_self_hosted_worker") .resolve("tool_ledger") .resolve(".tool-result-stale.tmp"); Files.write(stale, "partial".getBytes(java.nio.charset.StandardCharsets.UTF_8)); @@ -41,4 +42,35 @@ public void recoveryRemovesStaleTemporaryRecords() throws Exception { assertFalse(Files.exists(stale)); } + + @Test + public void sessionStoreIsolatesSessions() throws Exception { + Path workdir = Files.createTempDirectory("ark-java-store-"); + FileToolResultStore first = new FileToolResultStore(workdir.toString(), "session-a"); + FileToolResultStore second = new FileToolResultStore(workdir.toString(), "session-b"); + Map raw = new LinkedHashMap<>(); + raw.put("id", "event-1"); + raw.put("type", "agent.tool_use"); + raw.put("name", "bash"); + + first.begin("call-1", Event.fromMap(raw)); + + assertTrue(second.recover().getPending().isEmpty()); + assertTrue(Files.isDirectory(workdir + .resolve(".ma_self_hosted_worker") + .resolve("tool_ledger") + .resolve("session-a"))); + } + + @Test + public void sessionStoreSanitizesSessionId() throws Exception { + Path workdir = Files.createTempDirectory("ark-java-store-"); + new FileToolResultStore(workdir.toString(), "../../outside"); + Path base = workdir.resolve(".ma_self_hosted_worker").resolve("tool_ledger"); + + try (java.nio.file.DirectoryStream entries = Files.newDirectoryStream(base, "session-*")) { + assertTrue(entries.iterator().hasNext()); + } + assertFalse(Files.exists(workdir.resolve("outside"))); + } } diff --git a/src/test/java/com/volcengine/ark/runtime/selfhosted/InitializerTest.java b/src/test/java/com/volcengine/ark/runtime/selfhosted/InitializerTest.java index b6bcbfe..a51e603 100644 --- a/src/test/java/com/volcengine/ark/runtime/selfhosted/InitializerTest.java +++ b/src/test/java/com/volcengine/ark/runtime/selfhosted/InitializerTest.java @@ -4,6 +4,7 @@ package com.volcengine.ark.runtime.selfhosted; import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertFalse; import static org.junit.Assert.assertTrue; import java.io.ByteArrayOutputStream; @@ -71,6 +72,40 @@ public void zipArchiveEntryLimitIsEnforced() throws Exception { throw new AssertionError("expected archive entry limit failure"); } + @Test + public void cleanupRemovesOnlyInstalledSkills() throws Exception { + byte[] archive = zipWithTwoEntries(); + OkHttpClient http = new OkHttpClient.Builder().addInterceptor(chain -> { + Request request = chain.request(); + if (request.url().encodedPath().equals("/api/v3/skills/skill-1")) { + return response( + request, + MediaType.parse("application/json"), + ("{\"id\":\"skill-1\",\"object\":\"skill\",\"created_at\":1," + + "\"name\":\"demo\",\"latest_version\":\"1\"}") + .getBytes(java.nio.charset.StandardCharsets.UTF_8)); + } + return response(request, MediaType.parse("application/zip"), archive); + }).build(); + SelfHostedClient client = new SelfHostedClient.Builder() + .apiKey("test-key") + .baseUrl("https://ark.example.com/api/v3") + .httpClient(http) + .build(); + Path root = Files.createTempDirectory("ark-java-skill-cleanup-"); + Initializer initializer = new Initializer(client, new Initializer.Options(root.toString())); + Map raw = new LinkedHashMap<>(); + raw.put("skill_id", "skill-1"); + raw.put("version", "1"); + + initializer.installSkill("session-1", SkillRef.fromMap(raw)); + Path retained = Files.createDirectories(root.resolve("skills").resolve("retained")); + initializer.cleanup(); + + assertFalse(Files.exists(root.resolve("skills").resolve("demo"))); + assertTrue(Files.isDirectory(retained)); + } + @Test public void replaceSkillRollsBackOldVersionWhenCommitFails() throws Exception { Path root = Files.createTempDirectory("ark-java-skill-rollback-");