@mastra/redis 1.4.0-alpha.0 → 1.4.0-alpha.2

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
package/dist/index.js CHANGED
@@ -287,7 +287,7 @@ var StoreMemoryRedis = class extends MemoryStorage {
287
287
  });
288
288
  const updatedThread = {
289
289
  ...thread,
290
- title,
290
+ title: title ?? thread.title,
291
291
  metadata: {
292
292
  ...thread.metadata,
293
293
  ...metadata
@@ -408,16 +408,36 @@ var StoreMemoryRedis = class extends MemoryStorage {
408
408
  if (message.threadId) await this.client.set(getMessageIndexKey(messageId), message.threadId);
409
409
  return message.threadId || null;
410
410
  }
411
- async getIncludedMessages(include) {
411
+ /**
412
+ * Fetches the messages named by `include` together with their surrounding context.
413
+ *
414
+ * @param include - Message ids to pin, each with an optional before/after window.
415
+ * @param resourceId - When set, drops any pinned or context message owned by another
416
+ * resource so an id from another resource returns nothing.
417
+ */
418
+ async getIncludedMessages(include, resourceId) {
412
419
  if (!include?.length) return [];
413
420
  const messageIds = /* @__PURE__ */ new Set();
414
421
  const messageIdToThreadIds = {};
415
422
  for (const item of include) {
416
423
  const itemThreadId = await this.getThreadIdForMessage(item.id);
417
424
  if (!itemThreadId) continue;
425
+ const itemThreadMessagesKey = getThreadMessagesKey(itemThreadId);
426
+ if (resourceId !== void 0) {
427
+ const threadMessageIds = await this.client.zRange(itemThreadMessagesKey, 0, -1);
428
+ const threadMessages = (await this.client.mGet(threadMessageIds.map((id) => getMessageKey(itemThreadId, id)))).filter((data) => data !== null).map((data) => JSON.parse(data)).filter((message) => message.resourceId === resourceId);
429
+ const targetIndex = threadMessages.findIndex((message) => message.id === item.id);
430
+ if (targetIndex === -1) continue;
431
+ const start = Math.max(0, targetIndex - (item.withPreviousMessages ?? 0));
432
+ const end = Math.min(threadMessages.length, targetIndex + (item.withNextMessages ?? 0) + 1);
433
+ for (const message of threadMessages.slice(start, end)) {
434
+ messageIds.add(message.id);
435
+ messageIdToThreadIds[message.id] = itemThreadId;
436
+ }
437
+ continue;
438
+ }
418
439
  messageIds.add(item.id);
419
440
  messageIdToThreadIds[item.id] = itemThreadId;
420
- const itemThreadMessagesKey = getThreadMessagesKey(itemThreadId);
421
441
  const rank = await this.client.zRank(itemThreadMessagesKey, item.id);
422
442
  if (rank === null) continue;
423
443
  if (item.withPreviousMessages) {
@@ -434,7 +454,8 @@ var StoreMemoryRedis = class extends MemoryStorage {
434
454
  }
435
455
  if (messageIds.size === 0) return [];
436
456
  const keysToFetch = Array.from(messageIds).map((id) => getMessageKey(messageIdToThreadIds[id], id));
437
- return (await this.client.mGet(keysToFetch)).filter((data) => data !== null).map((data) => JSON.parse(data));
457
+ const includedMessages = (await this.client.mGet(keysToFetch)).filter((data) => data !== null).map((data) => JSON.parse(data));
458
+ return resourceId ? includedMessages.filter((message) => message.resourceId === resourceId) : includedMessages;
438
459
  }
439
460
  parseStoredMessage(storedMessage) {
440
461
  const defaultMessageContent = {
@@ -538,7 +559,7 @@ var StoreMemoryRedis = class extends MemoryStorage {
538
559
  hasMore: false
539
560
  };
540
561
  let includedMessages = [];
541
- if (include && include.length > 0) includedMessages = (await this.getIncludedMessages(include)).map(this.parseStoredMessage);
562
+ if (include && include.length > 0) includedMessages = (await this.getIncludedMessages(include, resourceId)).map(this.parseStoredMessage);
542
563
  if (perPage === 0 && include && include.length > 0) return {
543
564
  messages: new MessageList().add(includedMessages, "memory").get.all.db().sort((a, b) => {
544
565
  const aValue = getFieldValue(a);