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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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);
Expand All @@ -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<String> 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);
Expand All @@ -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)
Expand All @@ -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 {
Expand Down Expand Up @@ -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) {
Expand Down Expand Up @@ -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);
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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);
}

Expand Down Expand Up @@ -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<String, Object> asMap(Object value) {
return value instanceof Map ? (Map<String, Object>) value : Collections.<String, Object>emptyMap();
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -26,6 +27,7 @@ public class Initializer {

private final SelfHostedClient api;
private final Options options;
private final List<Path> installedSkillDirs = new ArrayList<>();

public Initializer(SelfHostedClient api, Options options) {
if (api == null) {
Expand Down Expand Up @@ -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));
Expand All @@ -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 {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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();
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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));
Expand All @@ -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<String, Object> 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<Path> entries = Files.newDirectoryStream(base, "session-*")) {
assertTrue(entries.iterator().hasNext());
}
assertFalse(Files.exists(workdir.resolve("outside")));
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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<String, Object> 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-");
Expand Down
Loading