@mastra/redis 1.3.1 → 1.4.0-alpha.1

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
@@ -240,6 +240,7 @@ var StoreMemoryRedis = class extends MemoryStorage {
240
240
  hasMore: perPageInput === false ? false : end < total
241
241
  };
242
242
  } catch (error) {
243
+ if (error instanceof MastraError && error.category === ErrorCategory.USER) throw error;
243
244
  const mastraError = new MastraError({
244
245
  id: createStorageErrorId("REDIS", "LIST_THREADS", "FAILED"),
245
246
  domain: ErrorDomain.STORAGE,
@@ -253,13 +254,7 @@ var StoreMemoryRedis = class extends MemoryStorage {
253
254
  }, error);
254
255
  this.logger.trackException(mastraError);
255
256
  this.logger.error(mastraError.toString());
256
- return {
257
- threads: [],
258
- total: 0,
259
- page,
260
- perPage: perPageForResponse,
261
- hasMore: false
262
- };
257
+ throw mastraError;
263
258
  }
264
259
  }
265
260
  async saveThread({ thread }) {
@@ -413,16 +408,36 @@ var StoreMemoryRedis = class extends MemoryStorage {
413
408
  if (message.threadId) await this.client.set(getMessageIndexKey(messageId), message.threadId);
414
409
  return message.threadId || null;
415
410
  }
416
- 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) {
417
419
  if (!include?.length) return [];
418
420
  const messageIds = /* @__PURE__ */ new Set();
419
421
  const messageIdToThreadIds = {};
420
422
  for (const item of include) {
421
423
  const itemThreadId = await this.getThreadIdForMessage(item.id);
422
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
+ }
423
439
  messageIds.add(item.id);
424
440
  messageIdToThreadIds[item.id] = itemThreadId;
425
- const itemThreadMessagesKey = getThreadMessagesKey(itemThreadId);
426
441
  const rank = await this.client.zRank(itemThreadMessagesKey, item.id);
427
442
  if (rank === null) continue;
428
443
  if (item.withPreviousMessages) {
@@ -439,7 +454,8 @@ var StoreMemoryRedis = class extends MemoryStorage {
439
454
  }
440
455
  if (messageIds.size === 0) return [];
441
456
  const keysToFetch = Array.from(messageIds).map((id) => getMessageKey(messageIdToThreadIds[id], id));
442
- 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;
443
459
  }
444
460
  parseStoredMessage(storedMessage) {
445
461
  const defaultMessageContent = {
@@ -543,7 +559,7 @@ var StoreMemoryRedis = class extends MemoryStorage {
543
559
  hasMore: false
544
560
  };
545
561
  let includedMessages = [];
546
- 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);
547
563
  if (perPage === 0 && include && include.length > 0) return {
548
564
  messages: new MessageList().add(includedMessages, "memory").get.all.db().sort((a, b) => {
549
565
  const aValue = getFieldValue(a);
@@ -613,6 +629,7 @@ var StoreMemoryRedis = class extends MemoryStorage {
613
629
  hasMore
614
630
  };
615
631
  } catch (error) {
632
+ if (error instanceof MastraError && error.category === ErrorCategory.USER) throw error;
616
633
  const mastraError = new MastraError({
617
634
  id: createStorageErrorId("REDIS", "LIST_MESSAGES", "FAILED"),
618
635
  domain: ErrorDomain.STORAGE,
@@ -624,13 +641,7 @@ var StoreMemoryRedis = class extends MemoryStorage {
624
641
  }, error);
625
642
  this.logger.error(mastraError.toString());
626
643
  this.logger.trackException(mastraError);
627
- return {
628
- messages: [],
629
- total: 0,
630
- page,
631
- perPage: perPageForResponse,
632
- hasMore: false
633
- };
644
+ throw mastraError;
634
645
  }
635
646
  }
636
647
  async getResourceById({ resourceId }) {