Skip to content
Open
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 @@ -696,7 +696,9 @@ public void getInterpreterBindings(NotebookSocket conn,
setting.getInterpreterInfos(), true));
}
}
conn.send(serializeMessage(new Message(OP.INTERPRETER_BINDINGS).put("interpreterBindings", settingList)));
conn.send(serializeMessage(new Message(OP.INTERPRETER_BINDINGS)
.put("noteId", noteId)
.put("interpreterBindings", settingList)));
return null;
});
}
Expand Down Expand Up @@ -734,7 +736,9 @@ public void saveInterpreterBindings(NotebookSocket conn, ServiceContext context,
});
if (permitted) {
conn.send(serializeMessage(
new Message(OP.INTERPRETER_BINDINGS).put("interpreterBindings", settingList)));
new Message(OP.INTERPRETER_BINDINGS)
.put("noteId", noteId)
.put("interpreterBindings", settingList)));
}
}

Expand Down Expand Up @@ -1669,7 +1673,9 @@ public void onSuccess(Revision revision, ServiceContext context) throws IOExcept

List<Revision> revisions = getNotebook().processNote(noteId,
note -> getNotebook().listRevisionHistory(noteId, note.getPath(), context.getAutheInfo()));
conn.send(serializeMessage(new Message(OP.LIST_REVISION_HISTORY).put("revisionList", revisions)));
conn.send(serializeMessage(new Message(OP.LIST_REVISION_HISTORY)
.put("noteId", noteId)
.put("revisionList", revisions)));
} else {
conn.send(serializeMessage(
new Message(OP.ERROR_INFO).put("info",
Expand All @@ -1689,7 +1695,9 @@ private void listRevisionHistory(NotebookSocket conn,
@Override
public void onSuccess(List<Revision> revisions, ServiceContext context) throws IOException {
super.onSuccess(revisions, context);
conn.send(serializeMessage(new Message(OP.LIST_REVISION_HISTORY).put("revisionList", revisions)));
conn.send(serializeMessage(new Message(OP.LIST_REVISION_HISTORY)
.put("noteId", noteId)
.put("revisionList", revisions)));
}
});
}
Expand All @@ -1705,7 +1713,9 @@ private void setNoteRevision(NotebookSocket conn,
public void onSuccess(Note note, ServiceContext context) throws IOException {
super.onSuccess(note, context);
Note reloadedNote = getNotebook().loadNoteFromRepo(noteId, context.getAutheInfo());
conn.send(serializeMessage(new Message(OP.SET_NOTE_REVISION).put("status", true)));
conn.send(serializeMessage(new Message(OP.SET_NOTE_REVISION)
.put("noteId", noteId)
.put("status", true)));
broadcastNote(reloadedNote);
}
});
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1067,15 +1067,17 @@ void getInterpreterBindingsRequiresReaderPermission() throws IOException {

ArgumentCaptor<String> response = ArgumentCaptor.forClass(String.class);
verify(socket).send(response.capture());
assertEquals(OP.AUTH_INFO, notebookServer.deserializeMessage(response.getValue()).op);
Message responseMessage = notebookServer.deserializeMessage(response.getValue());
assertEquals(OP.AUTH_INFO, responseMessage.op);

reset(socket);
setNotePermissions(noteId, "binding-owner", "binding-reader");
notebookServer.getInterpreterBindings(socket, serviceContext("binding-reader"), message);

verify(socket).send(response.capture());
assertEquals(OP.INTERPRETER_BINDINGS,
notebookServer.deserializeMessage(response.getValue()).op);
responseMessage = notebookServer.deserializeMessage(response.getValue());
assertEquals(OP.INTERPRETER_BINDINGS, responseMessage.op);
assertEquals(noteId, responseMessage.data.get("noteId"));
} finally {
notebook.removeNote(noteId, owner);
}
Expand All @@ -1100,7 +1102,8 @@ void saveInterpreterBindingsRequiresWriterPermission() throws IOException {
notebook.processNote(noteId, Note::getDefaultInterpreterGroup));
ArgumentCaptor<String> response = ArgumentCaptor.forClass(String.class);
verify(socket).send(response.capture());
assertEquals(OP.AUTH_INFO, notebookServer.deserializeMessage(response.getValue()).op);
Message responseMessage = notebookServer.deserializeMessage(response.getValue());
assertEquals(OP.AUTH_INFO, responseMessage.op);

reset(socket);
authorizationService.setWriters(noteId,
Expand All @@ -1110,8 +1113,9 @@ void saveInterpreterBindingsRequiresWriterPermission() throws IOException {
assertEquals(replacementGroup,
notebook.processNote(noteId, Note::getDefaultInterpreterGroup));
verify(socket).send(response.capture());
assertEquals(OP.INTERPRETER_BINDINGS,
notebookServer.deserializeMessage(response.getValue()).op);
responseMessage = notebookServer.deserializeMessage(response.getValue());
assertEquals(OP.INTERPRETER_BINDINGS, responseMessage.op);
assertEquals(noteId, responseMessage.data.get("noteId"));
} finally {
notebook.removeNote(noteId, owner);
}
Expand Down Expand Up @@ -1170,6 +1174,82 @@ void testNoteRevision() throws IOException {
}
}

@Test
void listRevisionHistoryIncludesNoteId() throws IOException {
String noteId = notebook.createNote("revision-list-note-id", anonymous);

try {
NotebookSocket socket = createWebSocket();
Message request = new Message(OP.LIST_REVISION_HISTORY)
.put("noteId", noteId);

notebookServer.onMessage(socket, request.toJson());

ArgumentCaptor<String> response = ArgumentCaptor.forClass(String.class);
verify(socket).send(response.capture());

Message responseMessage = notebookServer.deserializeMessage(response.getValue());
assertEquals(OP.LIST_REVISION_HISTORY, responseMessage.op);
assertEquals(noteId, responseMessage.data.get("noteId"));
} finally {
notebook.removeNote(noteId, anonymous);
}
}

@Test
void checkpointNoteIncludesNoteId() throws IOException {
String noteId = notebook.createNote("checkpoint-note-id", anonymous);

try {
NotebookSocket socket = createWebSocket();
Message request = new Message(OP.CHECKPOINT_NOTE)
.put("noteId", noteId)
.put("commitMessage", "checkpoint");

notebookServer.onMessage(socket, request.toJson());

ArgumentCaptor<String> response = ArgumentCaptor.forClass(String.class);
verify(socket).send(response.capture());

Message responseMessage = notebookServer.deserializeMessage(response.getValue());
assertEquals(OP.LIST_REVISION_HISTORY, responseMessage.op);
assertEquals(noteId, responseMessage.data.get("noteId"));
} finally {
notebook.removeNote(noteId, anonymous);
}
}

@Test
void setNoteRevisionIncludesNoteId() throws IOException {
String noteId = notebook.createNote("set-revision-note-id", anonymous);

try {
NotebookRepoWithVersionControl.Revision revision =
notebook.processNote(noteId,
note -> notebook.checkpointNote(
note.getId(),
note.getPath(),
"revision",
anonymous));

NotebookSocket socket = createWebSocket();
Message request = new Message(OP.SET_NOTE_REVISION)
.put("noteId", noteId)
.put("revisionId", revision.id);

notebookServer.onMessage(socket, request.toJson());

ArgumentCaptor<String> response = ArgumentCaptor.forClass(String.class);
verify(socket).send(response.capture());

Message responseMessage = notebookServer.deserializeMessage(response.getValue());
assertEquals(OP.SET_NOTE_REVISION, responseMessage.op);
assertEquals(noteId, responseMessage.data.get("noteId"));
} finally {
notebook.removeNote(noteId, anonymous);
}
}

private NotebookSocket createWebSocket() {
NotebookSocket sock = mock(NotebookSocket.class);
return sock;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@ export interface InterpreterItem {
}

export interface InterpreterBindings {
noteId: string;
interpreterBindings: InterpreterBindingItem[];
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -167,15 +167,16 @@ export interface ImportNoteReceived {

export interface ParagraphAdded {
index: number;
msgId?: string;
paragraph: ParagraphItem;
}

export interface SetNoteRevisionStatus {
noteId: string;
status: boolean;
}

export interface ListRevision {
noteId: string;
revisionList: RevisionListItem[];
}

Expand Down
18 changes: 18 additions & 0 deletions zeppelin-web-angular/projects/zeppelin-sdk/src/message.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -143,4 +143,22 @@ describe('Message.receive', () => {

expect(listener).toHaveBeenCalledWith(data);
});

it('passes the full message envelope with msgId', () => {
const message = new Message();
const listener = vi.fn();
const data = {};

message.receiveEnvelope(OP.NOTE).subscribe(listener);

const envelope = asReceivedMessage({
op: OP.NOTE,
msgId: 'note-request-1',
data
});

message.shortCircuit(envelope);

expect(listener).toHaveBeenCalledWith(envelope);
});
});
39 changes: 25 additions & 14 deletions zeppelin-web-angular/projects/zeppelin-sdk/src/message.ts
Original file line number Diff line number Diff line change
Expand Up @@ -175,21 +175,15 @@ export class Message {
}

receive<K extends keyof MessageReceiveDataTypeMap>(op: K): Observable<Record<K, MessageReceiveDataTypeMap[K]>[K]> {
const guard = getMessagePayloadGuard(op);

return this.received$.pipe(
filter(message => message.op === op),
filter(message => {
if (!guard || guard(message.data)) {
return true;
}
return this.receiveMessage(op).pipe(map(message => message.data)) as Observable<
Record<K, MessageReceiveDataTypeMap[K]>[K]
>;
}

// The payload can be large and carries note names, so log the OP alone.
console.warn(`Dropped WebSocket OP ${String(op)}: payload failed validation`);
return false;
}),
map(message => message.data)
) as Observable<Record<K, MessageReceiveDataTypeMap[K]>[K]>;
receiveEnvelope<K extends keyof MessageReceiveDataTypeMap>(
op: K
): Observable<WebSocketMessage<MessageReceiveDataTypeMap, K>> {
return this.receiveMessage(op) as Observable<WebSocketMessage<MessageReceiveDataTypeMap, K>>;
}

shortCircuit(message: WebSocketMessage<MessageReceiveDataTypeMap>) {
Expand Down Expand Up @@ -565,4 +559,21 @@ export class Message {
formName
});
}

private receiveMessage<K extends keyof MessageReceiveDataTypeMap>(op: K) {
const guard = getMessagePayloadGuard(op);

return this.received$.pipe(
filter(message => message.op === op),
filter(message => {
if (!guard || guard(message.data)) {
return true;
}

// The payload can be large and carries note names, so log the OP alone.
console.warn(`Dropped WebSocket OP ${String(op)}: payload failed validation`);
return false;
})
);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -13,9 +13,8 @@
import { Subject } from 'rxjs';
import { afterEach, describe, expect, it, vi } from 'vitest';

import { Message, OP, MessageReceiveDataTypeMap } from '@zeppelin/sdk';

import { MessageListener, MessageListenersManager } from './message-listener';
import { Message, OP, MessageReceiveDataTypeMap, type WebSocketMessage } from '@zeppelin/sdk';
import { MessageEnvelopeListener, MessageListener, MessageListenersManager } from './message-listener';

afterEach(() => {
vi.restoreAllMocks();
Expand Down Expand Up @@ -88,3 +87,36 @@ describe('MessageListener', () => {
expect(component.receivedData).toBe(data);
});
});

describe('MessageEnvelopeListener', () => {
it('passes the received message envelope with msgId to the handler', () => {
const received$ = new Subject<WebSocketMessage<MessageReceiveDataTypeMap, OP.NOTE>>();
const messageService = {
receiveEnvelope: vi.fn(() => received$.asObservable())
} as unknown as Message;

class TestComponent extends MessageListenersManager {
receivedMessage?: WebSocketMessage<MessageReceiveDataTypeMap, OP.NOTE>;

handleNote(message: WebSocketMessage<MessageReceiveDataTypeMap, OP.NOTE>): void {
this.receivedMessage = message;
}
}

const descriptor = Object.getOwnPropertyDescriptor(TestComponent.prototype, 'handleNote')!;

MessageEnvelopeListener(OP.NOTE)(TestComponent.prototype, 'handleNote', descriptor);

const component = new TestComponent(messageService);
const envelope: WebSocketMessage<MessageReceiveDataTypeMap, OP.NOTE> = {
op: OP.NOTE,
msgId: 'note-request-1',
data: {} as MessageReceiveDataTypeMap[OP.NOTE]
};

received$.next(envelope);

expect(component.receivedMessage).toBe(envelope);
expect(messageService.receiveEnvelope).toHaveBeenCalledWith(OP.NOTE);
});
});
Original file line number Diff line number Diff line change
Expand Up @@ -11,9 +11,9 @@
*/

import { Component, OnDestroy } from '@angular/core';
import { Subscriber } from 'rxjs';
import { Observable, Subscriber } from 'rxjs';

import { Message, MessageReceiveDataTypeMap, ReceiveArgumentsType } from '@zeppelin/sdk';
import { Message, MessageReceiveDataTypeMap } from '@zeppelin/sdk';

@Component({
template: '',
Expand All @@ -34,21 +34,26 @@ export class MessageListenersManager implements OnDestroy {
}
}

export function MessageListener<K extends keyof MessageReceiveDataTypeMap>(op: K) {
type ListenerArgumentsType<T> = T extends undefined ? () => void : (data: T) => void;

const createMessageListener = <K extends keyof MessageReceiveDataTypeMap, T>(
op: K,
receiver: (messageService: Message, op: K) => Observable<T>
) => {
return function (
target: MessageListenersManager,
propertyKey: string,
descriptor: TypedPropertyDescriptor<ReceiveArgumentsType<K>>
descriptor: TypedPropertyDescriptor<ListenerArgumentsType<T>>
) {
const oldValue = descriptor.value as ReceiveArgumentsType<K>;
const oldValue = descriptor.value as ListenerArgumentsType<T>;

const fn = function (this: MessageListenersManager) {
if (!this.__zeppelinMessageListeners$__) {
throw new Error('__zeppelinMessageListeners$__ is not defined');
}

this.__zeppelinMessageListeners$__.add(
this.messageService.receive(op).subscribe(data => {
receiver(this.messageService, op).subscribe(data => {
try {
// @ts-ignore
oldValue.apply(this, [data]);
Expand All @@ -68,4 +73,12 @@ export function MessageListener<K extends keyof MessageReceiveDataTypeMap>(op: K

return descriptor;
};
}
};

export const MessageListener = <K extends keyof MessageReceiveDataTypeMap>(op: K) => {
return createMessageListener(op, (messageService, targetOp) => messageService.receive(targetOp));
};

export const MessageEnvelopeListener = <K extends keyof MessageReceiveDataTypeMap>(op: K) => {
return createMessageListener(op, (messageService, targetOp) => messageService.receiveEnvelope(targetOp));
};
Loading
Loading