Skip to content

Commit 030b48c

Browse files
kaiban-distributed - change:Fixed DistributedStateMiddleware Bug to publish to global Redis pub/sub channel
1 parent 738cc45 commit 030b48c

5 files changed

Lines changed: 91 additions & 74 deletions

File tree

package-lock.json

Lines changed: 3 additions & 3 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

src/adapters/state/distributedMiddleware.ts

Lines changed: 22 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,5 @@
1-
import {
2-
IMessagingDriver,
3-
MessagePayload,
4-
} from "../../infrastructure/messaging/interfaces";
1+
import { Redis } from 'ioredis';
2+
import type { MessagePayload } from "../../infrastructure/messaging/interfaces";
53

64
const PII_DENYLIST: ReadonlySet<string> = new Set([
75
'email', 'name', 'phone', 'ip', 'password', 'token', 'secret', 'ssn', 'dob',
@@ -23,11 +21,11 @@ interface ZustandStore {
2321
}
2422

2523
export class DistributedStateMiddleware {
26-
private driver: IMessagingDriver;
24+
private redis: Redis;
2725
private channelName: string;
2826

29-
constructor(driver: IMessagingDriver, channelName = 'kaiban-state-events') {
30-
this.driver = driver;
27+
constructor(redisUrl: string, channelName = 'kaiban-state-events') {
28+
this.redis = new Redis(redisUrl, { lazyConnect: false });
3129
this.channelName = channelName;
3230
}
3331

@@ -46,7 +44,7 @@ export class DistributedStateMiddleware {
4644
};
4745

4846
try {
49-
await this.driver.publish(this.channelName, payload);
47+
await this.redis.publish(this.channelName, JSON.stringify(payload));
5048
} catch (err) {
5149
console.error("[DistributedStateMiddleware] Failed to publish state delta:", err);
5250
}
@@ -56,11 +54,21 @@ export class DistributedStateMiddleware {
5654
}
5755

5856
public async listen(onStateChange: (delta: Record<string, unknown>) => void): Promise<void> {
59-
await this.driver.subscribe(
60-
this.channelName,
61-
async (payload: MessagePayload) => {
62-
onStateChange(payload.data['stateUpdate'] as Record<string, unknown>);
63-
},
64-
);
57+
const sub = new Redis(this.redis.options);
58+
await sub.subscribe(this.channelName);
59+
sub.on('message', (channel, message) => {
60+
if (channel === this.channelName) {
61+
try {
62+
const payload = JSON.parse(message) as MessagePayload;
63+
onStateChange(payload.data['stateUpdate'] as Record<string, unknown>);
64+
} catch (e) {
65+
console.error("[DistributedStateMiddleware] Failed to parse message:", e);
66+
}
67+
}
68+
});
69+
}
70+
71+
public async disconnect(): Promise<void> {
72+
await this.redis.quit();
6573
}
6674
}

src/infrastructure/kaibanjs/kaiban-team-bridge.ts

Lines changed: 7 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
11
import { Team } from 'kaibanjs';
22
import type { Agent, Task } from 'kaibanjs';
33
import { DistributedStateMiddleware } from '../../adapters/state/distributedMiddleware';
4-
import type { IMessagingDriver } from '../messaging/interfaces';
54

65
export interface KaibanTeamConfig {
76
name: string;
@@ -17,7 +16,7 @@ export interface KaibanTeamConfig {
1716
* SocketGateway → kaiban-board in real-time.
1817
*
1918
* Usage:
20-
* const bridge = new KaibanTeamBridge({ name, agents, tasks }, driver);
19+
* const bridge = new KaibanTeamBridge({ name, agents, tasks }, 'redis://localhost:6379');
2120
* const result = await bridge.start({ topic: 'AI news' });
2221
*/
2322
export class KaibanTeamBridge {
@@ -26,11 +25,11 @@ export class KaibanTeamBridge {
2625

2726
constructor(
2827
config: KaibanTeamConfig,
29-
driver: IMessagingDriver,
28+
redisUrl: string,
3029
stateChannel = 'kaiban-state-events',
3130
) {
3231
this.team = new Team({ name: config.name, agents: config.agents, tasks: config.tasks, env: config.env ?? {} });
33-
this.middleware = new DistributedStateMiddleware(driver, stateChannel);
32+
this.middleware = new DistributedStateMiddleware(redisUrl, stateChannel);
3433

3534
const store = this.team.getStore() as unknown as { setState: (p: Record<string, unknown>) => void };
3635
this.middleware.attach(store);
@@ -50,4 +49,8 @@ export class KaibanTeamBridge {
5049
): () => void {
5150
return this.team.subscribeToChanges(listener, properties);
5251
}
52+
53+
async disconnect(): Promise<void> {
54+
await this.middleware.disconnect();
55+
}
5356
}

tests/unit/kaibanjs/kaiban-team-bridge.test.ts

Lines changed: 15 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,5 @@
11
import { describe, it, expect, vi, beforeEach } from 'vitest';
22
import { KaibanTeamBridge } from '../../../src/infrastructure/kaibanjs/kaiban-team-bridge';
3-
import type { IMessagingDriver } from '../../../src/infrastructure/messaging/interfaces';
43

54
const mockGetStore = vi.fn().mockReturnValue({
65
setState: vi.fn(),
@@ -18,49 +17,51 @@ vi.mock('kaibanjs', () => ({
1817
Task: vi.fn().mockImplementation(function (params: Record<string, unknown>) { return params; }),
1918
}));
2019

21-
function makeMockDriver(): IMessagingDriver {
22-
return {
23-
publish: vi.fn().mockResolvedValue(undefined),
24-
subscribe: vi.fn().mockResolvedValue(undefined),
25-
unsubscribe: vi.fn().mockResolvedValue(undefined),
26-
disconnect: vi.fn().mockResolvedValue(undefined),
27-
};
28-
}
20+
const mockRedis = {
21+
publish: vi.fn().mockResolvedValue(undefined),
22+
subscribe: vi.fn().mockResolvedValue(undefined),
23+
quit: vi.fn().mockResolvedValue(undefined),
24+
on: vi.fn(),
25+
};
26+
27+
vi.mock('ioredis', () => ({
28+
Redis: vi.fn().mockImplementation(function() { return mockRedis; }),
29+
}));
2930

3031
describe('KaibanTeamBridge', () => {
3132
beforeEach(() => { vi.clearAllMocks(); });
3233

3334
it('calls team.getStore() to attach DistributedStateMiddleware', () => {
3435
const bridge = new KaibanTeamBridge(
3536
{ name: 'Blog Team', agents: [], tasks: [] },
36-
makeMockDriver(),
37+
'redis://localhost:6379',
3738
);
3839
expect(bridge).toBeDefined();
3940
expect(mockGetStore).toHaveBeenCalledOnce();
4041
});
4142

4243
it('getTeam() returns the underlying KaibanJS Team', () => {
43-
const bridge = new KaibanTeamBridge({ name: 'T', agents: [], tasks: [] }, makeMockDriver());
44+
const bridge = new KaibanTeamBridge({ name: 'T', agents: [], tasks: [] }, 'redis://localhost:6379');
4445
const team = bridge.getTeam();
4546
expect(team).toBeDefined();
4647
expect(typeof team.start).toBe('function');
4748
});
4849

4950
it('start() delegates to team.start() with inputs', async () => {
50-
const bridge = new KaibanTeamBridge({ name: 'T', agents: [], tasks: [] }, makeMockDriver());
51+
const bridge = new KaibanTeamBridge({ name: 'T', agents: [], tasks: [] }, 'redis://localhost:6379');
5152
const result = await bridge.start({ topic: 'AI trends' });
5253
expect(mockStart).toHaveBeenCalledWith({ topic: 'AI trends' });
5354
expect(result.status).toBe('FINISHED');
5455
});
5556

5657
it('start() with no inputs passes empty object', async () => {
57-
const bridge = new KaibanTeamBridge({ name: 'T', agents: [], tasks: [] }, makeMockDriver());
58+
const bridge = new KaibanTeamBridge({ name: 'T', agents: [], tasks: [] }, 'redis://localhost:6379');
5859
await bridge.start();
5960
expect(mockStart).toHaveBeenCalledWith({});
6061
});
6162

6263
it('subscribeToChanges() sets up a store listener', () => {
63-
const bridge = new KaibanTeamBridge({ name: 'T', agents: [], tasks: [] }, makeMockDriver());
64+
const bridge = new KaibanTeamBridge({ name: 'T', agents: [], tasks: [] }, 'redis://localhost:6379');
6465
const listener = vi.fn();
6566
bridge.subscribeToChanges(listener, ['teamWorkflowStatus']);
6667
expect(mockGetStore).toHaveBeenCalled();
Lines changed: 44 additions & 39 deletions
Original file line numberDiff line numberDiff line change
@@ -1,15 +1,18 @@
11
import { describe, it, expect, vi } from "vitest";
22
import { DistributedStateMiddleware } from "../../../src/adapters/state/distributedMiddleware";
3-
import { IMessagingDriver, MessagePayload } from "../../../src/infrastructure/messaging/interfaces";
3+
import { MessagePayload } from "../../../src/infrastructure/messaging/interfaces";
44

5-
function makeMockDriver(): IMessagingDriver {
6-
return {
7-
publish: vi.fn().mockResolvedValue(undefined),
8-
subscribe: vi.fn().mockResolvedValue(undefined),
9-
unsubscribe: vi.fn().mockResolvedValue(undefined),
10-
disconnect: vi.fn().mockResolvedValue(undefined),
11-
};
12-
}
5+
const mockRedis = {
6+
publish: vi.fn().mockResolvedValue(undefined),
7+
subscribe: vi.fn(),
8+
quit: vi.fn().mockResolvedValue(undefined),
9+
on: vi.fn(),
10+
options: {},
11+
};
12+
13+
vi.mock('ioredis', () => ({
14+
Redis: vi.fn().mockImplementation(function() { return mockRedis; }),
15+
}));
1316

1417
interface MockStore {
1518
state: Record<string, unknown>;
@@ -24,65 +27,68 @@ function makeStore(initial: Record<string, unknown> = {}): MockStore {
2427
}
2528

2629
describe("DistributedStateMiddleware", () => {
30+
beforeEach(() => {
31+
vi.clearAllMocks();
32+
});
33+
2734
it("intercepts setState and publishes to the driver", async () => {
28-
const mockDriver = makeMockDriver();
29-
const mw = new DistributedStateMiddleware(mockDriver);
35+
const mw = new DistributedStateMiddleware('redis://localhost');
3036
const store = makeStore({ count: 0 });
3137
mw.attach(store);
3238
await store.setState({ count: 1 });
3339
expect(store.state['count']).toBe(1);
34-
expect(mockDriver.publish).toHaveBeenCalledWith("kaiban-state-events", expect.objectContaining({ data: { stateUpdate: { count: 1 } } }));
40+
expect(mockRedis.publish).toHaveBeenCalledWith(
41+
"kaiban-state-events",
42+
expect.stringContaining('"stateUpdate":{"count":1}')
43+
);
3544
});
3645

3746
it("sanitizeDelta strips PII keys", async () => {
38-
const mockDriver = makeMockDriver();
39-
const mw = new DistributedStateMiddleware(mockDriver);
47+
const mw = new DistributedStateMiddleware('redis://localhost');
4048
const store = makeStore();
4149
mw.attach(store);
4250
await store.setState({ count: 2, email: "x@y.com", token: "abc", password: "pw" });
43-
const call = (mockDriver.publish as ReturnType<typeof vi.fn>).mock.calls[0];
44-
const data = call[1].data.stateUpdate as Record<string, unknown>;
51+
const call = mockRedis.publish.mock.calls[0];
52+
const parsed = JSON.parse(call[1] as string) as MessagePayload;
53+
const data = parsed.data['stateUpdate'] as Record<string, unknown>;
4554
expect(data['count']).toBe(2);
4655
expect(data['email']).toBeUndefined();
4756
expect(data['token']).toBeUndefined();
4857
expect(data['password']).toBeUndefined();
4958
});
5059

5160
it("sanitizeDelta handles null partial (covers null branch)", async () => {
52-
const mockDriver = makeMockDriver();
53-
const mw = new DistributedStateMiddleware(mockDriver);
61+
const mw = new DistributedStateMiddleware('redis://localhost');
5462
const store = makeStore();
5563
mw.attach(store);
5664
// Trigger with null via coercion to cover `if (partial === null) return {}`
5765
await (store.setState as (p: unknown) => Promise<void>)(null);
58-
const call = (mockDriver.publish as ReturnType<typeof vi.fn>).mock.calls[0];
59-
expect(call[1].data.stateUpdate).toEqual({});
66+
const call = mockRedis.publish.mock.calls[0];
67+
const parsed = JSON.parse(call[1] as string) as MessagePayload;
68+
expect(parsed.data['stateUpdate']).toEqual({});
6069
});
6170

6271
it("listen() subscribes and delivers state deltas via callback", async () => {
63-
let capturedHandler!: (payload: MessagePayload) => Promise<void>;
64-
const mockDriver: IMessagingDriver = {
65-
publish: vi.fn().mockResolvedValue(undefined),
66-
subscribe: vi.fn((_q, handler) => { capturedHandler = handler; return Promise.resolve(); }),
67-
unsubscribe: vi.fn().mockResolvedValue(undefined),
68-
disconnect: vi.fn().mockResolvedValue(undefined),
69-
};
70-
const mw = new DistributedStateMiddleware(mockDriver);
72+
let capturedHandler!: (channel: string, message: string) => void;
73+
mockRedis.on.mockImplementation((event, handler) => {
74+
if (event === 'message') capturedHandler = handler;
75+
});
76+
77+
const mw = new DistributedStateMiddleware('redis://localhost');
7178
const onStateChange = vi.fn();
7279
await mw.listen(onStateChange);
73-
await capturedHandler({ taskId: "g", agentId: "system", timestamp: 0, data: { stateUpdate: { x: 1 } } });
80+
81+
// Simulate incoming Redis pub/sub message
82+
const msg = JSON.stringify({ taskId: "g", agentId: "system", timestamp: 0, data: { stateUpdate: { x: 1 } } });
83+
capturedHandler("kaiban-state-events", msg);
84+
7485
expect(onStateChange).toHaveBeenCalledWith({ x: 1 });
7586
});
7687

7788
it("publish error is caught and logged without throwing", async () => {
78-
const mockDriver: IMessagingDriver = {
79-
publish: vi.fn().mockRejectedValue(new Error("redis down")),
80-
subscribe: vi.fn().mockResolvedValue(undefined),
81-
unsubscribe: vi.fn().mockResolvedValue(undefined),
82-
disconnect: vi.fn().mockResolvedValue(undefined),
83-
};
89+
mockRedis.publish.mockRejectedValueOnce(new Error("redis down"));
8490
const errSpy = vi.spyOn(console, "error").mockImplementation(() => {});
85-
const mw = new DistributedStateMiddleware(mockDriver);
91+
const mw = new DistributedStateMiddleware('redis://localhost');
8692
const store = makeStore();
8793
mw.attach(store);
8894
await expect(store.setState({ x: 1 })).resolves.not.toThrow();
@@ -91,11 +97,10 @@ describe("DistributedStateMiddleware", () => {
9197
});
9298

9399
it("sanitizeDelta handles empty state update", async () => {
94-
const mockDriver = makeMockDriver();
95-
const mw = new DistributedStateMiddleware(mockDriver);
100+
const mw = new DistributedStateMiddleware('redis://localhost');
96101
const store = makeStore();
97102
mw.attach(store);
98103
await store.setState({});
99-
expect(mockDriver.publish).toHaveBeenCalledOnce();
104+
expect(mockRedis.publish).toHaveBeenCalledOnce();
100105
});
101106
});

0 commit comments

Comments
 (0)