fix: replace PostHog evaluate_flags with static config for OSS notices (#7185)

This commit is contained in:
Kartik
2026-09-02 18:00:01 +05:30
committed by GitHub
parent c33ca27f5e
commit 3cf41878ea
15 changed files with 643 additions and 468 deletions
+64 -29
View File
@@ -1,11 +1,9 @@
import * as crypto from "crypto";
import * as fs from "fs";
import * as os from "os";
import * as path from "path";
import {
POSTHOG_API_KEY,
captureNoticeEvent,
isTelemetryEnabled,
} from "./telemetry";
import { captureNoticeEvent, isTelemetryEnabled } from "./telemetry";
import bundledNoticeConfig from "./oss_notices_config.json";
import type { TelemetryInstance } from "./telemetry.types";
export const NOTICE_FLAG_KEY = "mem0-oss-notices";
@@ -21,7 +19,57 @@ export const PERFORMANCE_SLOW_QUERY_NOTICE_ID = "performance_slow_query";
export const NOTICE_CAP_LIMIT = 10;
export const NOTICE_CAP_WINDOW_MS = 7 * 24 * 60 * 60 * 1000;
export const NOTICE_FLAG_TIMEOUT_MS = 500;
export const POSTHOG_FLAGS_URL = "https://us.i.posthog.com/flags?v=2";
const REMOTE_CONFIG_URL =
"https://raw.githubusercontent.com/mem0ai/mem0/main/mem0/memory/oss_notices_config.json";
const CONFIG_TTL_MS = 3_600_000;
let _cachedConfig: Record<string, any> | null = null;
let _cachedConfigTs = 0;
let _configFetch: Promise<Record<string, any>> | null = null;
async function fetchNoticeConfig(
fetchImpl: typeof fetch,
timeoutMs: number,
): Promise<Record<string, any>> {
const now = Date.now();
if (_cachedConfig && now - _cachedConfigTs < CONFIG_TTL_MS) {
return _cachedConfig;
}
if (_configFetch) return _configFetch;
_configFetch = (async () => {
try {
const controller = new AbortController();
const timeout = setTimeout(() => controller.abort(), timeoutMs);
try {
const resp = await fetchImpl(REMOTE_CONFIG_URL, {
signal: controller.signal,
});
if (resp.ok) {
_cachedConfig = await resp.json();
}
} finally {
clearTimeout(timeout);
}
} catch {}
_cachedConfig ??= bundledNoticeConfig as Record<string, any>;
_cachedConfigTs = Date.now();
return _cachedConfig;
})();
try {
return await _configFetch;
} finally {
_configFetch = null;
}
}
function resetNoticeConfigCache(): void {
_cachedConfig = null;
_cachedConfigTs = 0;
_configFetch = null;
}
const DISPLAYED_VARIANT = "displayed";
const HOLDOUT_VARIANT = "holdout";
const LOG_LINE_NOTICE_TYPE = "log_line";
@@ -255,37 +303,23 @@ export async function evaluateNoticeFlag(
const timeoutMs = options.timeoutMs ?? NOTICE_FLAG_TIMEOUT_MS;
const fetchImpl = options.fetchImpl ?? fetch;
const controller = new AbortController();
const timeout = setTimeout(() => controller.abort(), timeoutMs);
try {
const response = await fetchImpl(POSTHOG_FLAGS_URL, {
method: "POST",
headers: { "Content-Type": "application/json" },
body: JSON.stringify({
api_key: POSTHOG_API_KEY,
distinct_id: distinctId,
}),
signal: controller.signal,
});
if (!response.ok) return null;
const data: any = await response.json();
const flag = data?.flags?.[NOTICE_FLAG_KEY];
if (!flag || flag.enabled === false) return null;
const variant = typeof flag.variant === "string" ? flag.variant : null;
if (!variant) return null;
const config = await fetchNoticeConfig(fetchImpl, timeoutMs);
const split =
typeof config.variant_split === "number" ? config.variant_split : 0.5;
const hashKey = `${NOTICE_FLAG_KEY}.${distinctId}variant`;
const hash = crypto.createHash("sha1").update(hashKey).digest("hex");
const bucket = parseInt(hash.substring(0, 15), 16) / 0xfffffffffffffff;
const variant = bucket < split ? DISPLAYED_VARIANT : HOLDOUT_VARIANT;
return {
variant,
payload: parsePayload(flag.metadata?.payload),
flag,
payload: { notices: config.notices ?? {} },
flag: { key: NOTICE_FLAG_KEY, enabled: true, variant },
};
} catch {
return null;
} finally {
clearTimeout(timeout);
}
}
@@ -1408,6 +1442,7 @@ export const __noticeTestHooks = {
appendNoticeCapEvent,
emitNoticeDisplayed,
evaluateNoticeFlag,
resetNoticeConfigCache,
displayFirstRunNotice,
displayDecayUsageNotice,
displayTemporalUsageNotice,
@@ -0,0 +1,44 @@
{
"version": 1,
"variant_split": 0.5,
"notices": {
"first_run": {
"copy": "",
"enabled": false,
"notice_type": "log_line"
},
"temporal_stub": {
"copy": "",
"enabled": false,
"notice_type": "error"
},
"temporal_usage": {
"copy": "",
"enabled": false,
"notice_type": "log_line"
},
"decay_stub": {
"copy": "",
"enabled": false,
"notice_type": "error"
},
"decay_usage": {
"copy": "",
"enabled": false,
"notice_type": "log_line"
},
"scale_threshold": {
"copies": {
"memory_count": "",
"top_k": ""
},
"enabled": false,
"notice_type": "log_line"
},
"performance_slow_query": {
"copy": "",
"enabled": false,
"notice_type": "log_line"
}
}
}
@@ -57,34 +57,25 @@ function decayPayload(overrides: Record<string, any> = {}) {
function createFetchMock(options: {
variant?: string;
payload?: unknown;
failFlags?: boolean;
flagEnabled?: boolean;
payload?: Record<string, any>;
failConfig?: boolean;
}) {
const calls: any[] = [];
const variantSplit = options.variant === "holdout" ? 0.0 : 1.0;
const payload = options.payload ?? decayPayload();
const fetchMock = jest.fn(async (url: string | URL, init?: RequestInit) => {
const target = String(url);
if (target.includes("/flags")) {
if (options.failFlags) {
throw new Error("flag evaluation failed");
if (target.includes("raw.githubusercontent.com")) {
if (options.failConfig) {
throw new Error("config fetch failed");
}
return {
ok: true,
json: jest.fn().mockResolvedValue({
flags: {
"mem0-oss-notices": {
key: "mem0-oss-notices",
enabled: options.flagEnabled ?? true,
variant: options.variant ?? "displayed",
metadata: {
payload:
options.payload === undefined
? JSON.stringify(decayPayload())
: options.payload,
},
},
},
version: 1,
variant_split: variantSplit,
...payload,
}),
};
}
@@ -222,32 +213,10 @@ describe("Node OSS decay feature error notice", () => {
expect(notices[0].properties.bypass_reason).toBeUndefined();
});
it("uses plain error for unknown future variants and emits not_displayed", async () => {
const { fetchMock, calls } = createFetchMock({ variant: "silent" });
global.fetch = fetchMock as any;
const memory = await createMemory();
await expect(memory.updateProject({ decay: true })).rejects.toThrow(
PLAIN_DECAY_ERROR,
);
const notices = noticeEvents(calls);
expect(notices).toHaveLength(1);
expect(notices[0].properties).toEqual(
expect.objectContaining({
notice_id: "decay_stub",
variant: "silent",
displayed: false,
bypass_reason: "not_displayed",
payload: DECAY_COPY,
}),
);
});
it("uses plain error for disabled payload and emits payload_disabled", async () => {
const { fetchMock, calls } = createFetchMock({
variant: "holdout",
payload: JSON.stringify(decayPayload({ enabled: false })),
payload: decayPayload({ enabled: false }),
});
global.fetch = fetchMock as any;
const memory = await createMemory();
@@ -270,17 +239,8 @@ describe("Node OSS decay feature error notice", () => {
});
it.each([
[
"missing config",
JSON.stringify({ notices: {} }),
"missing_notice_config",
],
[
"missing copy",
JSON.stringify(decayPayload({ copy: "" })),
"missing_copy",
],
["malformed payload", "{not-json", "missing_notice_config"],
["missing config", { notices: {} }, "missing_notice_config"],
["missing copy", decayPayload({ copy: "" }), "missing_copy"],
])("uses plain error for %s", async (_label, payload, bypassReason) => {
const { fetchMock, calls } = createFetchMock({
variant: "displayed",
@@ -304,11 +264,8 @@ describe("Node OSS decay feature error notice", () => {
);
});
it("uses plain error and emits no event when the blunt flag is disabled", async () => {
const { fetchMock, calls } = createFetchMock({
variant: "displayed",
flagEnabled: false,
});
it("uses plain error when config fetch fails (bundled config fallback)", async () => {
const { fetchMock, calls } = createFetchMock({ failConfig: true });
global.fetch = fetchMock as any;
const memory = await createMemory();
@@ -316,19 +273,14 @@ describe("Node OSS decay feature error notice", () => {
PLAIN_DECAY_ERROR,
);
expect(noticeEvents(calls)).toHaveLength(0);
});
it("uses plain error and emits no event when PostHog fails", async () => {
const { fetchMock, calls } = createFetchMock({ failFlags: true });
global.fetch = fetchMock as any;
const memory = await createMemory();
await expect(memory.updateProject({ decay: true })).rejects.toThrow(
PLAIN_DECAY_ERROR,
const notices = noticeEvents(calls);
expect(notices).toHaveLength(1);
expect(notices[0].properties).toEqual(
expect.objectContaining({
notice_id: "decay_stub",
displayed: false,
}),
);
expect(noticeEvents(calls)).toHaveLength(0);
});
it("uses plain error and skips flag evaluation when telemetry is off", async () => {
@@ -361,7 +313,9 @@ describe("Node OSS decay feature error notice", () => {
);
expect(
fetchMock.mock.calls.filter(([url]) => String(url).includes("/flags")),
fetchMock.mock.calls.filter(([url]) =>
String(url).includes("raw.githubusercontent.com"),
),
).toHaveLength(0);
expect(noticeEvents(calls)).toHaveLength(0);
});
@@ -81,33 +81,25 @@ function decayUsagePayload(overrides: Record<string, any> = {}) {
function createFetchMock(options: {
variant?: string;
payload?: unknown;
failFlags?: boolean;
flagEnabled?: boolean;
payload?: Record<string, any>;
failConfig?: boolean;
}) {
const calls: any[] = [];
const variantSplit = options.variant === "holdout" ? 0.0 : 1.0;
const payload = options.payload ?? decayUsagePayload();
const fetchMock = jest.fn(async (url: string | URL, init?: RequestInit) => {
const target = String(url);
if (target.includes("/flags")) {
if (options.failFlags) {
throw new Error("flag evaluation failed");
if (target.includes("raw.githubusercontent.com")) {
if (options.failConfig) {
throw new Error("config fetch failed");
}
const payload =
options.payload === undefined
? JSON.stringify(decayUsagePayload())
: options.payload;
return {
ok: true,
json: jest.fn().mockResolvedValue({
flags: {
"mem0-oss-notices": {
key: "mem0-oss-notices",
enabled: options.flagEnabled ?? true,
variant: options.variant ?? "displayed",
metadata: { payload },
},
},
version: 1,
variant_split: variantSplit,
...payload,
}),
};
}
@@ -134,9 +126,10 @@ function noticeEvents(calls: any[]) {
return calls.filter((call) => call.event === "mem0.notice_displayed");
}
function flagRequestCount(fetchMock: jest.Mock) {
return fetchMock.mock.calls.filter(([url]) => String(url).includes("/flags"))
.length;
function configRequestCount(fetchMock: jest.Mock) {
return fetchMock.mock.calls.filter(([url]) =>
String(url).includes("raw.githubusercontent.com"),
).length;
}
async function createMemory() {
@@ -226,7 +219,7 @@ describe("Node OSS decay usage notice", () => {
await memory.delete(id);
}
expect(flagRequestCount(fetchMock)).toBe(0);
expect(configRequestCount(fetchMock)).toBe(0);
expect(noticeEvents(calls)).toHaveLength(0);
expect(readConfig().notice_state?.decay_usage).toBeUndefined();
expect(stderrSpy).not.toHaveBeenCalledWith(
@@ -279,7 +272,7 @@ describe("Node OSS decay usage notice", () => {
await expect(memory.delete("memory-id")).rejects.toThrow("delete failed");
expect(flagRequestCount(fetchMock)).toBe(0);
expect(configRequestCount(fetchMock)).toBe(0);
expect(noticeEvents(calls)).toHaveLength(0);
expect(readConfig().notice_state?.decay_usage).toBeUndefined();
});
@@ -317,7 +310,7 @@ describe("Node OSS decay usage notice", () => {
await memory.deleteAll({ userId: "empty-delete-all-user" });
expect(flagRequestCount(fetchMock)).toBe(0);
expect(configRequestCount(fetchMock)).toBe(0);
expect(noticeEvents(calls)).toHaveLength(0);
expect(readConfig().notice_state?.decay_usage).toBeUndefined();
});
@@ -332,7 +325,7 @@ describe("Node OSS decay usage notice", () => {
"At least one filter is required",
);
expect(flagRequestCount(fetchMock)).toBe(0);
expect(configRequestCount(fetchMock)).toBe(0);
expect(noticeEvents(calls)).toHaveLength(0);
expect(readConfig().notice_state?.decay_usage).toBeUndefined();
});
@@ -365,27 +358,21 @@ describe("Node OSS decay usage notice", () => {
it.each([
[
"disabled payload",
JSON.stringify(decayUsagePayload({ enabled: false })),
decayUsagePayload({ enabled: false }),
"payload_disabled",
"payload_disabled",
],
[
"missing config",
JSON.stringify({ notices: {} }),
"missing_notice_config",
undefined,
],
["missing config", { notices: {} }, "missing_notice_config", undefined],
[
"missing copy",
JSON.stringify({
{
notices: {
decay_usage: { enabled: true, notice_type: "log_line" },
},
}),
},
"missing_copy",
undefined,
],
["malformed payload", "{not-json", "missing_notice_config", undefined],
])(
"stays silent and emits safe bypass for %s",
async (_label, payload, bypassReason, disabledReason) => {
@@ -416,33 +403,23 @@ describe("Node OSS decay usage notice", () => {
},
);
it("does not emit or consume cap when the blunt flag is disabled", async () => {
it("emits not-displayed event when config fetch fails (bundled config fallback)", async () => {
consumeFirstRun();
const { fetchMock, calls } = createFetchMock({
variant: "displayed",
flagEnabled: false,
});
const { fetchMock, calls } = createFetchMock({ failConfig: true });
global.fetch = fetchMock as any;
const memory = await createMemory();
await addMemories(memory, "decay-flag-disabled-user", 3);
await addMemories(memory, "decay-config-failure-user", 3);
await memory.deleteAll({ userId: "decay-flag-disabled-user" });
await memory.deleteAll({ userId: "decay-config-failure-user" });
expect(noticeEvents(calls)).toHaveLength(0);
expect(readConfig().notice_state?.decay_usage).toBeUndefined();
});
it("does not emit or consume cap when PostHog fails", async () => {
consumeFirstRun();
const { fetchMock, calls } = createFetchMock({ failFlags: true });
global.fetch = fetchMock as any;
const memory = await createMemory();
await addMemories(memory, "decay-posthog-failure-user", 3);
await memory.deleteAll({ userId: "decay-posthog-failure-user" });
expect(noticeEvents(calls)).toHaveLength(0);
expect(readConfig().notice_state?.decay_usage).toBeUndefined();
const notices = noticeEvents(calls);
expect(notices).toHaveLength(1);
expect(notices[0].properties).toEqual(
expect.objectContaining({
notice_id: "decay_usage",
displayed: false,
}),
);
});
it("skips flag evaluation, event emission, and state writes when telemetry is off", async () => {
@@ -484,7 +461,7 @@ describe("Node OSS decay usage notice", () => {
await memory.deleteAll({ userId: "decay-cap-user" });
expect(flagRequestCount(fetchMock)).toBe(0);
expect(configRequestCount(fetchMock)).toBe(0);
expect(noticeEvents(calls)).toHaveLength(0);
expect(readConfig().notice_state.decay_usage.events).toHaveLength(10);
});
@@ -493,7 +470,7 @@ describe("Node OSS decay usage notice", () => {
consumeFirstRun();
const { fetchMock, calls } = createFetchMock({
variant: "displayed",
payload: JSON.stringify({
payload: {
notices: {
first_run: {
enabled: true,
@@ -506,7 +483,7 @@ describe("Node OSS decay usage notice", () => {
copy: DECAY_USAGE_COPY,
},
},
}),
},
});
global.fetch = fetchMock as any;
const memory = await createMemory();
+23 -19
View File
@@ -57,29 +57,24 @@ function firstRunPayload(overrides: Record<string, any> = {}) {
function createFetchMock(options: {
variant?: string;
payload?: Record<string, any>;
failFlags?: boolean;
failConfig?: boolean;
}) {
const calls: any[] = [];
const variantSplit = options.variant === "holdout" ? 0.0 : 1.0;
const payload = options.payload ?? firstRunPayload();
const fetchMock = jest.fn(async (url: string | URL, init?: RequestInit) => {
const target = String(url);
if (target.includes("/flags")) {
if (options.failFlags) {
throw new Error("flag evaluation failed");
if (target.includes("raw.githubusercontent.com")) {
if (options.failConfig) {
throw new Error("config fetch failed");
}
return {
ok: true,
json: jest.fn().mockResolvedValue({
flags: {
"mem0-oss-notices": {
key: "mem0-oss-notices",
enabled: true,
variant: options.variant ?? "displayed",
metadata: {
payload: JSON.stringify(options.payload ?? firstRunPayload()),
},
},
},
version: 1,
variant_split: variantSplit,
...payload,
}),
};
}
@@ -376,23 +371,32 @@ describe("Node OSS first-run notice", () => {
);
expect(
fetchMock.mock.calls.filter(([url]) => String(url).includes("/flags")),
fetchMock.mock.calls.filter(([url]) =>
String(url).includes("raw.githubusercontent.com"),
),
).toHaveLength(0);
expect(noticeEvents(calls)).toHaveLength(0);
});
it("does not break the successful operation when flag evaluation fails", async () => {
const { fetchMock, calls } = createFetchMock({ failFlags: true });
it("does not break the successful operation when config fetch fails", async () => {
const { fetchMock, calls } = createFetchMock({ failConfig: true });
global.fetch = fetchMock as any;
const memory = await createMemory();
const result = await memory.add("Flag failure content", {
const result = await memory.add("Config failure content", {
userId: "first-run-failure",
infer: false,
});
expect(result.results).toHaveLength(1);
expect(noticeEvents(calls)).toHaveLength(0);
const notices = noticeEvents(calls);
expect(notices).toHaveLength(1);
expect(notices[0].properties).toEqual(
expect.objectContaining({
notice_id: "first_run",
displayed: false,
}),
);
});
it("can trigger from get() when get is the first successful public call", async () => {
@@ -138,16 +138,49 @@ describe("Node OSS notice foundation", () => {
expect(fs.existsSync(notices.getMem0ConfigPath())).toBe(false);
});
it("returns null when PostHog flag evaluation fails", async () => {
it("falls back to bundled config when remote fetch fails", async () => {
const notices = await import("../src/utils/notices");
const fetchMock = jest.fn().mockRejectedValue(new Error("network down"));
await expect(
notices.evaluateNoticeFlag("notice-test-user", { fetchImpl: fetchMock }),
).resolves.toBeNull();
const firstResult = await notices.evaluateNoticeFlag("notice-test-user", {
fetchImpl: fetchMock,
});
const secondResult = await notices.evaluateNoticeFlag("notice-test-user", {
fetchImpl: fetchMock,
});
expect(firstResult).not.toBeNull();
expect(secondResult?.variant).toBe(firstResult?.variant);
expect(fetchMock).toHaveBeenCalledTimes(1);
});
it("returns null when PostHog flag evaluation times out", async () => {
it("coalesces concurrent remote config fetches", async () => {
const notices = await import("../src/utils/notices");
const resolvers: Array<(response: any) => void> = [];
const fetchMock = jest.fn(
() => new Promise((resolve) => resolvers.push(resolve)),
);
const evaluations = [
notices.evaluateNoticeFlag("user-0", { fetchImpl: fetchMock as any }),
notices.evaluateNoticeFlag("user-1", { fetchImpl: fetchMock as any }),
];
resolvers.forEach((resolve) =>
resolve({
ok: true,
json: jest.fn().mockResolvedValue({
version: 1,
variant_split: 0.5,
notices: {},
}),
}),
);
await Promise.all(evaluations);
expect(fetchMock).toHaveBeenCalledTimes(1);
});
it("falls back to bundled config when remote fetch times out", async () => {
const notices = await import("../src/utils/notices");
const fetchMock = jest.fn(
(_url: string | URL | Request, init?: RequestInit) =>
@@ -158,34 +191,28 @@ describe("Node OSS notice foundation", () => {
}),
);
await expect(
notices.evaluateNoticeFlag("notice-test-user", {
fetchImpl: fetchMock as any,
timeoutMs: 1,
}),
).resolves.toBeNull();
const startedAt = Date.now();
const result = await notices.evaluateNoticeFlag("notice-test-user", {
fetchImpl: fetchMock as any,
timeoutMs: 10,
});
expect(result).not.toBeNull();
expect(result?.variant).toMatch(/^(displayed|holdout)$/);
expect(Date.now() - startedAt).toBeLessThan(500);
});
it("parses displayed variant and JSON payload from PostHog flags response", async () => {
it("parses variant and notices from static config response", async () => {
const notices = await import("../src/utils/notices");
const payload = {
notices: {
foundation_notice: {
enabled: true,
notice_type: "log_line",
copy: "Foundation notice",
},
},
};
const fetchMock = jest.fn().mockResolvedValue({
ok: true,
json: jest.fn().mockResolvedValue({
flags: {
"mem0-oss-notices": {
key: "mem0-oss-notices",
version: 1,
variant_split: 1.0,
notices: {
foundation_notice: {
enabled: true,
variant: "displayed",
metadata: { payload: JSON.stringify(payload) },
notice_type: "log_line",
copy: "Foundation notice",
},
},
}),
@@ -203,6 +230,27 @@ describe("Node OSS notice foundation", () => {
expect(parsed.config?.copy).toBe("Foundation notice");
});
it.each([
["user-0", "holdout"],
["user-1", "displayed"],
])("preserves the PostHog cohort for %s", async (distinctId, expected) => {
const notices = await import("../src/utils/notices");
const fetchMock = jest.fn().mockResolvedValue({
ok: true,
json: jest.fn().mockResolvedValue({
version: 1,
variant_split: 0.5,
notices: {},
}),
});
const result = await notices.evaluateNoticeFlag(distinctId, {
fetchImpl: fetchMock,
});
expect(result?.variant).toBe(expected);
});
it("emits mem0.notice_displayed with sample_rate=1", async () => {
const fetchMock = jest.fn().mockResolvedValue({
ok: true,
@@ -81,33 +81,25 @@ function performancePayload(overrides: Record<string, any> = {}) {
function createFetchMock(options: {
variant?: string;
payload?: unknown;
failFlags?: boolean;
flagEnabled?: boolean;
payload?: Record<string, any>;
failConfig?: boolean;
}) {
const calls: any[] = [];
const variantSplit = options.variant === "holdout" ? 0.0 : 1.0;
const payload = options.payload ?? performancePayload();
const fetchMock = jest.fn(async (url: string | URL, init?: RequestInit) => {
const target = String(url);
if (target.includes("/flags")) {
if (options.failFlags) {
throw new Error("flag evaluation failed");
if (target.includes("raw.githubusercontent.com")) {
if (options.failConfig) {
throw new Error("config fetch failed");
}
const payload =
options.payload === undefined
? JSON.stringify(performancePayload())
: options.payload;
return {
ok: true,
json: jest.fn().mockResolvedValue({
flags: {
"mem0-oss-notices": {
key: "mem0-oss-notices",
enabled: options.flagEnabled ?? true,
variant: options.variant ?? "displayed",
metadata: { payload },
},
},
version: 1,
variant_split: variantSplit,
...payload,
}),
};
}
@@ -140,9 +132,10 @@ function performanceEvents(calls: any[]) {
);
}
function flagRequestCount(fetchMock: jest.Mock) {
return fetchMock.mock.calls.filter(([url]) => String(url).includes("/flags"))
.length;
function configRequestCount(fetchMock: jest.Mock) {
return fetchMock.mock.calls.filter(([url]) =>
String(url).includes("raw.githubusercontent.com"),
).length;
}
function mockSearchElapsed(elapsedMs: number) {
@@ -294,20 +287,11 @@ describe("Node OSS performance slow query notice", () => {
it.each([
[
"disabled payload",
JSON.stringify(performancePayload({ enabled: false })),
performancePayload({ enabled: false }),
"payload_disabled",
],
[
"missing config",
JSON.stringify({ notices: {} }),
"missing_notice_config",
],
[
"missing copy",
JSON.stringify(performancePayload({ copy: "" })),
"missing_copy",
],
["malformed payload", "{not-json", "missing_notice_config"],
["missing config", { notices: {} }, "missing_notice_config"],
["missing copy", performancePayload({ copy: "" }), "missing_copy"],
])("is silent and safe for %s", async (_label, payload, bypassReason) => {
consumeFirstRun();
const { fetchMock, calls } = createFetchMock({
@@ -347,7 +331,7 @@ describe("Node OSS performance slow query notice", () => {
topK: 3,
});
expect(flagRequestCount(fetchMock)).toBe(0);
expect(configRequestCount(fetchMock)).toBe(0);
expect(performanceEvents(calls)).toHaveLength(0);
expect(readConfig().notice_state?.performance_slow_query).toBeUndefined();
});
@@ -366,17 +350,14 @@ describe("Node OSS performance slow query notice", () => {
}),
).rejects.toThrow("filters must contain");
expect(flagRequestCount(fetchMock)).toBe(0);
expect(configRequestCount(fetchMock)).toBe(0);
expect(performanceEvents(calls)).toHaveLength(0);
expect(readConfig().notice_state?.performance_slow_query).toBeUndefined();
});
it("does not consume cap or emit when the blunt flag is disabled", async () => {
it("emits not-displayed when config fetch fails (bundled config fallback)", async () => {
consumeFirstRun();
const { fetchMock, calls } = createFetchMock({
variant: "displayed",
flagEnabled: false,
});
const { fetchMock, calls } = createFetchMock({ failConfig: true });
global.fetch = fetchMock as any;
const memory = await createMemory();
await addSeed(memory);
@@ -387,25 +368,14 @@ describe("Node OSS performance slow query notice", () => {
topK: 3,
});
expect(performanceEvents(calls)).toHaveLength(0);
expect(readConfig().notice_state?.performance_slow_query).toBeUndefined();
});
it("does not consume cap or emit when PostHog fails", async () => {
consumeFirstRun();
const { fetchMock, calls } = createFetchMock({ failFlags: true });
global.fetch = fetchMock as any;
const memory = await createMemory();
await addSeed(memory);
mockSearchElapsed(2345);
await memory.search("favorite drink", {
filters: { user_id: "performance-user" },
topK: 3,
});
expect(performanceEvents(calls)).toHaveLength(0);
expect(readConfig().notice_state?.performance_slow_query).toBeUndefined();
const notices = performanceEvents(calls);
expect(notices).toHaveLength(1);
expect(notices[0].properties).toEqual(
expect.objectContaining({
notice_id: "performance_slow_query",
displayed: false,
}),
);
});
it("does nothing when telemetry is off", async () => {
@@ -444,7 +414,7 @@ describe("Node OSS performance slow query notice", () => {
}
expect(performanceEvents(calls)).toHaveLength(10);
expect(flagRequestCount(fetchMock)).toBe(10);
expect(configRequestCount(fetchMock)).toBeLessThanOrEqual(1);
expect(stderrSpy).toHaveBeenCalledTimes(10);
expect(
readConfig().notice_state.performance_slow_query.events,
@@ -87,33 +87,25 @@ function scalePayload(overrides: Record<string, any> = {}) {
function createFetchMock(options: {
variant?: string;
payload?: unknown;
failFlags?: boolean;
flagEnabled?: boolean;
payload?: Record<string, any>;
failConfig?: boolean;
}) {
const calls: any[] = [];
const variantSplit = options.variant === "holdout" ? 0.0 : 1.0;
const payload = options.payload ?? scalePayload();
const fetchMock = jest.fn(async (url: string | URL, init?: RequestInit) => {
const target = String(url);
if (target.includes("/flags")) {
if (options.failFlags) {
throw new Error("flag evaluation failed");
if (target.includes("raw.githubusercontent.com")) {
if (options.failConfig) {
throw new Error("config fetch failed");
}
const payload =
options.payload === undefined
? JSON.stringify(scalePayload())
: options.payload;
return {
ok: true,
json: jest.fn().mockResolvedValue({
flags: {
"mem0-oss-notices": {
key: "mem0-oss-notices",
enabled: options.flagEnabled ?? true,
variant: options.variant ?? "displayed",
metadata: { payload },
},
},
version: 1,
variant_split: variantSplit,
...payload,
}),
};
}
@@ -146,9 +138,10 @@ function scaleEvents(calls: any[]) {
);
}
function flagRequestCount(fetchMock: jest.Mock) {
return fetchMock.mock.calls.filter(([url]) => String(url).includes("/flags"))
.length;
function configRequestCount(fetchMock: jest.Mock) {
return fetchMock.mock.calls.filter(([url]) =>
String(url).includes("raw.githubusercontent.com"),
).length;
}
async function createMemory() {
@@ -290,14 +283,14 @@ describe("Node OSS scale threshold notice", () => {
topK: 49,
});
expect(flagRequestCount(fetchMock)).toBe(0);
expect(configRequestCount(fetchMock)).toBe(0);
expect(scaleEvents(calls)).toHaveLength(0);
expect(readConfig().notice_state?.scale_threshold).toBeUndefined();
});
it("marks memory-count threshold evaluated before PostHog display", async () => {
it("marks memory-count threshold evaluated even when config fetch fails", async () => {
consumeFirstRun();
const { fetchMock, calls } = createFetchMock({ failFlags: true });
const { fetchMock, calls } = createFetchMock({ failConfig: true });
global.fetch = fetchMock as any;
const memory = await createMemory();
(memory as any).vectorStore.count = jest
@@ -314,7 +307,14 @@ describe("Node OSS scale threshold notice", () => {
readConfig().notice_state.scale_threshold
.memory_count_threshold_evaluated,
).toBe(true);
expect(scaleEvents(calls)).toHaveLength(0);
const notices = scaleEvents(calls);
expect(notices).toHaveLength(1);
expect(notices[0].properties).toEqual(
expect.objectContaining({
notice_id: "scale_threshold",
displayed: false,
}),
);
});
it("does not count provider memories once threshold was evaluated", async () => {
@@ -332,7 +332,7 @@ describe("Node OSS scale threshold notice", () => {
});
expect((memory as any).vectorStore.count).not.toHaveBeenCalled();
expect(flagRequestCount(fetchMock)).toBe(0);
expect(configRequestCount(fetchMock)).toBe(0);
});
it("throttles under-threshold provider counts", async () => {
@@ -354,14 +354,14 @@ describe("Node OSS scale threshold notice", () => {
});
expect((memory as any).vectorStore.count).toHaveBeenCalledTimes(1);
expect(flagRequestCount(fetchMock)).toBe(0);
expect(configRequestCount(fetchMock)).toBe(0);
});
it("records holdout and disabled variants silently", async () => {
consumeFirstRun();
const { fetchMock, calls } = createFetchMock({
variant: "holdout",
payload: JSON.stringify(scalePayload({ enabled: false })),
payload: scalePayload({ enabled: false }),
});
global.fetch = fetchMock as any;
const memory = await createMemory();
@@ -389,7 +389,7 @@ describe("Node OSS scale threshold notice", () => {
delete (payload.notices.scale_threshold as Record<string, any>).enabled;
const { fetchMock, calls } = createFetchMock({
variant: "displayed",
payload: JSON.stringify(payload),
payload,
});
global.fetch = fetchMock as any;
const memory = await createMemory();
@@ -425,7 +425,7 @@ describe("Node OSS scale threshold notice", () => {
}
expect(scaleEvents(calls)).toHaveLength(10);
expect(flagRequestCount(fetchMock)).toBe(10);
expect(configRequestCount(fetchMock)).toBeLessThanOrEqual(1);
expect(stderrSpy).toHaveBeenCalledTimes(10);
expect(readConfig().notice_state.scale_threshold.events).toHaveLength(10);
});
@@ -456,7 +456,7 @@ describe("Node OSS scale threshold notice", () => {
topK: 50,
});
expect(flagRequestCount(fetchMock)).toBe(0);
expect(configRequestCount(fetchMock)).toBe(0);
expect(scaleEvents(calls)).toHaveLength(0);
expect(readConfig().notice_state).toBeUndefined();
});
@@ -61,34 +61,25 @@ function temporalPayload(overrides: Record<string, any> = {}) {
function createFetchMock(options: {
variant?: string;
payload?: unknown;
failFlags?: boolean;
flagEnabled?: boolean;
payload?: Record<string, any>;
failConfig?: boolean;
}) {
const calls: any[] = [];
const variantSplit = options.variant === "holdout" ? 0.0 : 1.0;
const payload = options.payload ?? temporalPayload();
const fetchMock = jest.fn(async (url: string | URL, init?: RequestInit) => {
const target = String(url);
if (target.includes("/flags")) {
if (options.failFlags) {
throw new Error("flag evaluation failed");
if (target.includes("raw.githubusercontent.com")) {
if (options.failConfig) {
throw new Error("config fetch failed");
}
return {
ok: true,
json: jest.fn().mockResolvedValue({
flags: {
"mem0-oss-notices": {
key: "mem0-oss-notices",
enabled: options.flagEnabled ?? true,
variant: options.variant ?? "displayed",
metadata: {
payload:
options.payload === undefined
? JSON.stringify(temporalPayload())
: options.payload,
},
},
},
version: 1,
variant_split: variantSplit,
...payload,
}),
};
}
@@ -265,37 +256,12 @@ describe("Node OSS temporal feature error notice", () => {
},
);
it("uses plain error for unknown future variants and emits not_displayed", async () => {
const { fetchMock, calls } = createFetchMock({ variant: "silent" });
global.fetch = fetchMock as any;
const memory = await createMemory();
await expect(
memory.add("Temporal add", {
userId: "temporal-user",
timestamp: 1778112000,
}),
).rejects.toThrow(PLAIN_TIMESTAMP_ERROR);
const notices = noticeEvents(calls);
expect(notices).toHaveLength(1);
expect(notices[0].properties).toEqual(
expect.objectContaining({
notice_id: "temporal_stub",
variant: "silent",
displayed: false,
bypass_reason: "not_displayed",
payload: TEMPORAL_COPY,
}),
);
});
it("treats missing enabled as enabled for feature-error payloads", async () => {
const payload = temporalPayload();
delete (payload.notices.temporal_stub as Record<string, any>).enabled;
const { fetchMock, calls } = createFetchMock({
variant: "displayed",
payload: JSON.stringify(payload),
payload,
});
global.fetch = fetchMock as any;
const memory = await createMemory();
@@ -319,7 +285,7 @@ describe("Node OSS temporal feature error notice", () => {
it("uses timestamp plain error for disabled payload and emits payload_disabled", async () => {
const { fetchMock, calls } = createFetchMock({
variant: "displayed",
payload: JSON.stringify(temporalPayload({ enabled: false })),
payload: temporalPayload({ enabled: false }),
});
global.fetch = fetchMock as any;
const memory = await createMemory();
@@ -347,17 +313,8 @@ describe("Node OSS temporal feature error notice", () => {
});
it.each([
[
"missing config",
JSON.stringify({ notices: {} }),
"missing_notice_config",
],
[
"missing copy",
JSON.stringify(temporalPayload({ copy: "" })),
"missing_copy",
],
["malformed payload", "{not-json", "missing_notice_config"],
["missing config", { notices: {} }, "missing_notice_config"],
["missing copy", temporalPayload({ copy: "" }), "missing_copy"],
])(
"uses referenceDate plain error for %s",
async (_label, payload, bypassReason) => {
@@ -389,11 +346,8 @@ describe("Node OSS temporal feature error notice", () => {
},
);
it("uses plain error and emits no event when the blunt flag is disabled", async () => {
const { fetchMock, calls } = createFetchMock({
variant: "displayed",
flagEnabled: false,
});
it("uses plain error when config fetch fails (bundled config fallback)", async () => {
const { fetchMock, calls } = createFetchMock({ failConfig: true });
global.fetch = fetchMock as any;
const memory = await createMemory();
@@ -404,22 +358,14 @@ describe("Node OSS temporal feature error notice", () => {
}),
).rejects.toThrow(PLAIN_TIMESTAMP_ERROR);
expect(noticeEvents(calls)).toHaveLength(0);
});
it("uses plain error and emits no event when PostHog fails", async () => {
const { fetchMock, calls } = createFetchMock({ failFlags: true });
global.fetch = fetchMock as any;
const memory = await createMemory();
await expect(
memory.search("Temporal search", {
filters: { user_id: "temporal-user" },
referenceDate: "2026-05-06",
const notices = noticeEvents(calls);
expect(notices).toHaveLength(1);
expect(notices[0].properties).toEqual(
expect.objectContaining({
notice_id: "temporal_stub",
displayed: false,
}),
).rejects.toThrow(PLAIN_REFERENCE_DATE_ERROR);
expect(noticeEvents(calls)).toHaveLength(0);
);
});
it("uses plain error and skips flag evaluation when telemetry is off", async () => {
@@ -81,33 +81,25 @@ function temporalUsagePayload(overrides: Record<string, any> = {}) {
function createFetchMock(options: {
variant?: string;
payload?: unknown;
failFlags?: boolean;
flagEnabled?: boolean;
payload?: Record<string, any>;
failConfig?: boolean;
}) {
const calls: any[] = [];
const variantSplit = options.variant === "holdout" ? 0.0 : 1.0;
const payload = options.payload ?? temporalUsagePayload();
const fetchMock = jest.fn(async (url: string | URL, init?: RequestInit) => {
const target = String(url);
if (target.includes("/flags")) {
if (options.failFlags) {
throw new Error("flag evaluation failed");
if (target.includes("raw.githubusercontent.com")) {
if (options.failConfig) {
throw new Error("config fetch failed");
}
const payload =
options.payload === undefined
? JSON.stringify(temporalUsagePayload())
: options.payload;
return {
ok: true,
json: jest.fn().mockResolvedValue({
flags: {
"mem0-oss-notices": {
key: "mem0-oss-notices",
enabled: options.flagEnabled ?? true,
variant: options.variant ?? "displayed",
metadata: { payload },
},
},
version: 1,
variant_split: variantSplit,
...payload,
}),
};
}
@@ -140,9 +132,10 @@ function temporalUsageEvents(calls: any[]) {
);
}
function flagRequestCount(fetchMock: jest.Mock) {
return fetchMock.mock.calls.filter(([url]) => String(url).includes("/flags"))
.length;
function configRequestCount(fetchMock: jest.Mock) {
return fetchMock.mock.calls.filter(([url]) =>
String(url).includes("raw.githubusercontent.com"),
).length;
}
async function createMemory() {
@@ -320,7 +313,7 @@ describe("Node OSS temporal usage notice", () => {
topK: 3,
});
expect(flagRequestCount(fetchMock)).toBe(0);
expect(configRequestCount(fetchMock)).toBe(0);
expect(temporalUsageEvents(calls)).toHaveLength(0);
expect(readConfig().notice_state?.temporal_usage).toBeUndefined();
});
@@ -377,20 +370,11 @@ describe("Node OSS temporal usage notice", () => {
it.each([
[
"disabled payload",
JSON.stringify(temporalUsagePayload({ enabled: false })),
temporalUsagePayload({ enabled: false }),
"payload_disabled",
],
[
"missing config",
JSON.stringify({ notices: {} }),
"missing_notice_config",
],
[
"missing copy",
JSON.stringify(temporalUsagePayload({ copy: "" })),
"missing_copy",
],
["malformed payload", "{not-json", "missing_notice_config"],
["missing config", { notices: {} }, "missing_notice_config"],
["missing copy", temporalUsagePayload({ copy: "" }), "missing_copy"],
])("is silent and safe for %s", async (_label, payload, bypassReason) => {
consumeFirstRun();
const { fetchMock, calls } = createFetchMock({
@@ -418,12 +402,9 @@ describe("Node OSS temporal usage notice", () => {
);
});
it("does not consume cap or emit when the blunt flag is disabled", async () => {
it("emits not-displayed when config fetch fails (bundled config fallback)", async () => {
consumeFirstRun();
const { fetchMock, calls } = createFetchMock({
variant: "displayed",
flagEnabled: false,
});
const { fetchMock, calls } = createFetchMock({ failConfig: true });
global.fetch = fetchMock as any;
const memory = await createMemory();
@@ -433,24 +414,14 @@ describe("Node OSS temporal usage notice", () => {
metadata: { event_date: "2025-04-09" },
});
expect(temporalUsageEvents(calls)).toHaveLength(0);
expect(readConfig().notice_state?.temporal_usage).toBeUndefined();
});
it("does not consume cap or emit when PostHog fails", async () => {
consumeFirstRun();
const { fetchMock, calls } = createFetchMock({ failFlags: true });
global.fetch = fetchMock as any;
const memory = await createMemory();
await memory.add("Temporal metadata memory", {
userId: "temporal-user",
infer: false,
metadata: { event_date: "2025-04-09" },
});
expect(temporalUsageEvents(calls)).toHaveLength(0);
expect(readConfig().notice_state?.temporal_usage).toBeUndefined();
const notices = temporalUsageEvents(calls);
expect(notices).toHaveLength(1);
expect(notices[0].properties).toEqual(
expect.objectContaining({
notice_id: "temporal_usage",
displayed: false,
}),
);
});
it("does nothing when telemetry is off", async () => {
@@ -487,7 +458,7 @@ describe("Node OSS temporal usage notice", () => {
}
expect(temporalUsageEvents(calls)).toHaveLength(10);
expect(flagRequestCount(fetchMock)).toBe(10);
expect(configRequestCount(fetchMock)).toBeLessThanOrEqual(1);
expect(readConfig().notice_state.temporal_usage.events).toHaveLength(10);
expect(stderrSpy).toHaveBeenCalledTimes(10);
});
+97 -12
View File
@@ -1,16 +1,101 @@
import asyncio
import hashlib
import json
import logging
import os
import re
import sys
import threading
import time
import urllib.error
import urllib.request
from datetime import date, datetime, timedelta, timezone
from typing import Any, Dict, Optional, Tuple
from mem0.memory import telemetry as telemetry_module
from mem0.memory.setup import _load_config, _write_config
_logger = logging.getLogger(__name__)
FLAG_KEY = "mem0-oss-notices"
_REMOTE_CONFIG_URL = (
"https://raw.githubusercontent.com/mem0ai/mem0/main/mem0/memory/oss_notices_config.json"
)
_BUNDLED_CONFIG_PATH = os.path.join(os.path.dirname(__file__), "oss_notices_config.json")
_CONFIG_TTL_SECONDS = 3600
_CONFIG_FETCH_TIMEOUT_SECONDS = telemetry_module.FEATURE_FLAGS_REQUEST_TIMEOUT_SECONDS
_cached_config = None
_cached_config_ts = 0.0
_config_fetch_lock = threading.Lock()
class StaticFlagResult:
"""Drop-in replacement for PostHog's evaluate_flags return value."""
def __init__(self, variant, payload):
self._variant = variant
self._payload = payload
def get_flag(self, key):
return self._variant
def get_flag_payload(self, key):
return self._payload
def _get_event_properties(self):
return {
f"$feature/{FLAG_KEY}": self._variant,
"$active_feature_flags": [FLAG_KEY],
}
def _load_bundled_config():
try:
with open(_BUNDLED_CONFIG_PATH) as f:
return json.load(f)
except Exception:
return {"version": 1, "variant_split": 0.5, "notices": {}}
def _fetch_remote_config():
try:
req = urllib.request.Request(_REMOTE_CONFIG_URL, headers={"User-Agent": "mem0-oss"})
with urllib.request.urlopen(req, timeout=_CONFIG_FETCH_TIMEOUT_SECONDS) as resp:
return json.loads(resp.read())
except Exception:
_logger.debug("Failed to fetch remote notice config, using bundled fallback")
return None
def _get_notice_config():
global _cached_config, _cached_config_ts
now = time.monotonic()
if _cached_config is not None and (now - _cached_config_ts) < _CONFIG_TTL_SECONDS:
return _cached_config
with _config_fetch_lock:
if _cached_config is not None and (now - _cached_config_ts) < _CONFIG_TTL_SECONDS:
return _cached_config
remote = _fetch_remote_config()
if remote is not None and isinstance(remote, dict) and "notices" in remote:
_cached_config = remote
else:
_cached_config = _load_bundled_config()
_cached_config_ts = time.monotonic()
return _cached_config
def _evaluate_notice_flags(user_id):
config = _get_notice_config()
split = config.get("variant_split", 0.5)
hash_key = f"{FLAG_KEY}.{user_id}variant"
bucket = int(hashlib.sha1(hash_key.encode()).hexdigest()[:15], 16) / float(0xFFFFFFFFFFFFFFF)
variant = DISPLAYED_VARIANT if bucket < split else HOLDOUT_VARIANT
payload = {"notices": config.get("notices", {})}
return StaticFlagResult(variant, payload)
NOTICE_ID = "first_run"
TEMPORAL_FEATURE_NOTICE_ID = "temporal_stub"
TEMPORAL_USAGE_NOTICE_ID = "temporal_usage"
@@ -88,10 +173,10 @@ def display_first_run_notice(memory_instance, sync_type: str, trigger_function:
variant = None
try:
telemetry = telemetry_module._get_oss_telemetry()
if telemetry is None or telemetry.posthog is None or not telemetry.user_id:
if telemetry is None or not telemetry.user_id:
return
flags = telemetry.posthog.evaluate_flags(telemetry.user_id, flag_keys=[FLAG_KEY])
flags = _evaluate_notice_flags(telemetry.user_id)
variant = flags.get_flag(FLAG_KEY)
_update_first_run_variant(variant)
@@ -168,10 +253,10 @@ def display_temporal_usage_notice(
try:
telemetry = telemetry_module._get_oss_telemetry()
if telemetry is None or telemetry.posthog is None or not telemetry.user_id:
if telemetry is None or not telemetry.user_id:
return
flags = telemetry.posthog.evaluate_flags(telemetry.user_id, flag_keys=[FLAG_KEY])
flags = _evaluate_notice_flags(telemetry.user_id)
variant = flags.get_flag(FLAG_KEY)
if variant in (None, False):
return
@@ -302,10 +387,10 @@ def display_decay_usage_notice(
try:
telemetry = telemetry_module._get_oss_telemetry()
if telemetry is None or telemetry.posthog is None or not telemetry.user_id:
if telemetry is None or not telemetry.user_id:
return
flags = telemetry.posthog.evaluate_flags(telemetry.user_id, flag_keys=[FLAG_KEY])
flags = _evaluate_notice_flags(telemetry.user_id)
variant = flags.get_flag(FLAG_KEY)
if variant in (None, False):
return
@@ -480,10 +565,10 @@ def display_scale_threshold_notice(
try:
telemetry = telemetry_module._get_oss_telemetry()
if telemetry is None or telemetry.posthog is None or not telemetry.user_id:
if telemetry is None or not telemetry.user_id:
return
flags = telemetry.posthog.evaluate_flags(telemetry.user_id, flag_keys=[FLAG_KEY])
flags = _evaluate_notice_flags(telemetry.user_id)
variant = flags.get_flag(FLAG_KEY)
if variant in (None, False):
return
@@ -596,10 +681,10 @@ def display_performance_slow_query_notice(
try:
telemetry = telemetry_module._get_oss_telemetry()
if telemetry is None or telemetry.posthog is None or not telemetry.user_id:
if telemetry is None or not telemetry.user_id:
return
flags = telemetry.posthog.evaluate_flags(telemetry.user_id, flag_keys=[FLAG_KEY])
flags = _evaluate_notice_flags(telemetry.user_id)
variant = flags.get_flag(FLAG_KEY)
if variant in (None, False):
return
@@ -749,10 +834,10 @@ def _get_feature_error_message(
try:
telemetry = telemetry_module._get_oss_telemetry()
if telemetry is None or telemetry.posthog is None or not telemetry.user_id:
if telemetry is None or not telemetry.user_id:
return plain_error
flags = telemetry.posthog.evaluate_flags(telemetry.user_id, flag_keys=[FLAG_KEY])
flags = _evaluate_notice_flags(telemetry.user_id)
variant = flags.get_flag(FLAG_KEY)
if variant in (None, False):
return plain_error
+44
View File
@@ -0,0 +1,44 @@
{
"version": 1,
"variant_split": 0.5,
"notices": {
"first_run": {
"copy": "",
"enabled": false,
"notice_type": "log_line"
},
"temporal_stub": {
"copy": "",
"enabled": false,
"notice_type": "error"
},
"temporal_usage": {
"copy": "",
"enabled": false,
"notice_type": "log_line"
},
"decay_stub": {
"copy": "",
"enabled": false,
"notice_type": "error"
},
"decay_usage": {
"copy": "",
"enabled": false,
"notice_type": "log_line"
},
"scale_threshold": {
"copies": {
"memory_count": "",
"top_k": ""
},
"enabled": false,
"notice_type": "log_line"
},
"performance_slow_query": {
"copy": "",
"enabled": false,
"notice_type": "log_line"
}
}
}
+2
View File
@@ -92,10 +92,12 @@ pythonpath = ["."]
[tool.hatch.build]
include = [
"mem0/**/*.py",
"mem0/memory/oss_notices_config.json",
]
exclude = [
"**/*",
"!mem0/**/*.py",
"!mem0/memory/oss_notices_config.json",
]
[tool.hatch.build.targets.wheel]
+113 -21
View File
@@ -56,6 +56,7 @@ def notice_harness(monkeypatch):
config = {}
telemetry = MagicMock()
telemetry.user_id = "oss-user"
evaluate_mock = MagicMock()
def write_config(updated):
saved = deepcopy(updated)
@@ -66,13 +67,15 @@ def notice_harness(monkeypatch):
monkeypatch.setattr(notices, "_write_config", write_config)
monkeypatch.setattr(notices.telemetry_module, "MEM0_TELEMETRY", True)
monkeypatch.setattr(notices.telemetry_module, "_get_oss_telemetry", lambda: telemetry)
monkeypatch.setattr(notices, "_evaluate_notice_flags", evaluate_mock)
telemetry._evaluate_mock = evaluate_mock
return config, telemetry
def configure_flag(telemetry, variant, payload):
flags = FakeFlags(variant, payload)
telemetry.posthog.evaluate_flags.return_value = flags
telemetry._evaluate_mock.return_value = flags
return flags
@@ -176,7 +179,7 @@ def test_displayed_notice_logs_once_and_captures_event(notice_harness, capsys):
config, telemetry, flags = display_notice(notice_harness)
assert capsys.readouterr().err == "Mem0 OSS notice\n"
telemetry.posthog.evaluate_flags.assert_called_once_with("oss-user", flag_keys=[notices.FLAG_KEY])
telemetry._evaluate_mock.assert_called_once_with("oss-user")
telemetry.capture_event.assert_called_once()
event_name, props = telemetry.capture_event.call_args.args
assert event_name == notices.NOTICE_EVENT
@@ -254,7 +257,7 @@ def test_telemetry_disabled_does_not_touch_posthog_or_state(monkeypatch, capsys)
def test_posthog_failure_is_silent_and_consumes_first_run(notice_harness, capsys):
config, telemetry = notice_harness
telemetry.posthog.evaluate_flags.side_effect = RuntimeError("network unavailable")
telemetry._evaluate_mock.side_effect = RuntimeError("network unavailable")
notices.display_first_run_notice(MagicMock(), "sync", "add")
@@ -265,7 +268,7 @@ def test_posthog_failure_is_silent_and_consumes_first_run(notice_harness, capsys
def test_public_add_succeeds_when_first_run_flag_eval_fails(notice_harness):
_, telemetry = notice_harness
telemetry.posthog.evaluate_flags.side_effect = RuntimeError("network unavailable")
telemetry._evaluate_mock.side_effect = RuntimeError("network unavailable")
memory = Memory.__new__(Memory)
memory.config = SimpleNamespace(llm=SimpleNamespace(config={}))
memory._add_to_vector_store = MagicMock(return_value=[{"event": "ADD", "memory": "likes tea"}])
@@ -277,7 +280,7 @@ def test_public_add_succeeds_when_first_run_flag_eval_fails(notice_harness):
def test_public_search_succeeds_when_first_run_flag_eval_fails(notice_harness, monkeypatch):
_, telemetry = notice_harness
telemetry.posthog.evaluate_flags.side_effect = RuntimeError("network unavailable")
telemetry._evaluate_mock.side_effect = RuntimeError("network unavailable")
monkeypatch.setattr(memory_main, "capture_event", MagicMock())
memory = Memory.__new__(Memory)
memory.api_version = "v1.1"
@@ -325,7 +328,7 @@ def test_temporal_feature_displayed_returns_payload_copy_and_captures_event(noti
assert message == "Temporal CTA"
assert capsys.readouterr().err == ""
telemetry.posthog.evaluate_flags.assert_called_once_with("oss-user", flag_keys=[notices.FLAG_KEY])
telemetry._evaluate_mock.assert_called_once_with("oss-user")
telemetry.capture_event.assert_called_once()
event_name, props = telemetry.capture_event.call_args.args
assert event_name == notices.NOTICE_EVENT
@@ -413,7 +416,7 @@ def test_temporal_feature_telemetry_disabled_does_not_touch_posthog(monkeypatch,
def test_temporal_feature_posthog_failure_returns_plain_error(notice_harness, capsys):
_, telemetry = notice_harness
telemetry.posthog.evaluate_flags.side_effect = RuntimeError("network unavailable")
telemetry._evaluate_mock.side_effect = RuntimeError("network unavailable")
message = notices.get_temporal_feature_error_message("sync", "add", "timestamp")
@@ -429,14 +432,14 @@ def test_temporal_feature_cap_blocks_repeated_posthog_evaluation(notice_harness)
for _ in range(notices.FEATURE_ERROR_CAP):
assert notices.get_temporal_feature_error_message("sync", "add", "timestamp") == "Temporal CTA"
assert telemetry.posthog.evaluate_flags.call_count == notices.FEATURE_ERROR_CAP
assert telemetry._evaluate_mock.call_count == notices.FEATURE_ERROR_CAP
assert telemetry.capture_event.call_count == notices.FEATURE_ERROR_CAP
assert len(config["notice_state"]["temporal_stub"]["events"]) == notices.FEATURE_ERROR_CAP
message = notices.get_temporal_feature_error_message("sync", "add", "timestamp")
assert message == notices.TEMPORAL_FEATURE_ERROR_MESSAGES["timestamp"]
assert telemetry.posthog.evaluate_flags.call_count == notices.FEATURE_ERROR_CAP
assert telemetry._evaluate_mock.call_count == notices.FEATURE_ERROR_CAP
assert telemetry.capture_event.call_count == notices.FEATURE_ERROR_CAP
@@ -465,7 +468,7 @@ def test_decay_feature_displayed_returns_payload_copy_and_captures_event(notice_
assert message == "Decay CTA"
assert capsys.readouterr().err == ""
telemetry.posthog.evaluate_flags.assert_called_once_with("oss-user", flag_keys=[notices.FLAG_KEY])
telemetry._evaluate_mock.assert_called_once_with("oss-user")
telemetry.capture_event.assert_called_once()
event_name, props = telemetry.capture_event.call_args.args
assert event_name == notices.NOTICE_EVENT
@@ -552,7 +555,7 @@ def test_decay_feature_telemetry_disabled_does_not_touch_posthog(monkeypatch, ca
def test_decay_feature_posthog_failure_returns_plain_error(notice_harness, capsys):
_, telemetry = notice_harness
telemetry.posthog.evaluate_flags.side_effect = RuntimeError("network unavailable")
telemetry._evaluate_mock.side_effect = RuntimeError("network unavailable")
message = notices.get_decay_feature_error_message("sync", "project.update", "decay")
@@ -582,14 +585,14 @@ def test_decay_feature_cap_blocks_repeated_posthog_evaluation(notice_harness):
for _ in range(notices.FEATURE_ERROR_CAP):
assert notices.get_decay_feature_error_message("sync", "project.update", "decay") == "Decay CTA"
assert telemetry.posthog.evaluate_flags.call_count == notices.FEATURE_ERROR_CAP
assert telemetry._evaluate_mock.call_count == notices.FEATURE_ERROR_CAP
assert telemetry.capture_event.call_count == notices.FEATURE_ERROR_CAP
assert len(config["notice_state"]["decay_stub"]["events"]) == notices.FEATURE_ERROR_CAP
message = notices.get_decay_feature_error_message("sync", "project.update", "decay")
assert message == notices.DECAY_FEATURE_ERROR_MESSAGE
assert telemetry.posthog.evaluate_flags.call_count == notices.FEATURE_ERROR_CAP
assert telemetry._evaluate_mock.call_count == notices.FEATURE_ERROR_CAP
assert telemetry.capture_event.call_count == notices.FEATURE_ERROR_CAP
@@ -670,7 +673,7 @@ def test_temporal_usage_displayed_logs_and_captures_event(notice_harness, capsys
notices.display_temporal_usage_notice(MagicMock(), "sync", "search", "query", "relative_phrase")
assert capsys.readouterr().err == "Temporal usage CTA\n"
telemetry.posthog.evaluate_flags.assert_called_once_with("oss-user", flag_keys=[notices.FLAG_KEY])
telemetry._evaluate_mock.assert_called_once_with("oss-user")
telemetry.capture_event.assert_called_once()
event_name, props = telemetry.capture_event.call_args.args
assert event_name == notices.NOTICE_EVENT
@@ -771,7 +774,7 @@ def test_temporal_usage_cap_blocks_before_posthog_eval(notice_harness, capsys):
notices.display_temporal_usage_notice(MagicMock(), "sync", "search", "query", "relative_phrase")
assert capsys.readouterr().err == "Temporal usage CTA\n" * notices.TEMPORAL_USAGE_CAP
assert telemetry.posthog.evaluate_flags.call_count == notices.TEMPORAL_USAGE_CAP
assert telemetry._evaluate_mock.call_count == notices.TEMPORAL_USAGE_CAP
assert telemetry.capture_event.call_count == notices.TEMPORAL_USAGE_CAP
assert len(config["notice_state"]["temporal_usage"]["events"]) == notices.TEMPORAL_USAGE_CAP
@@ -869,7 +872,7 @@ def test_decay_usage_displayed_logs_and_captures_event(notice_harness, capsys):
)
assert capsys.readouterr().err == "Decay usage CTA\n"
telemetry.posthog.evaluate_flags.assert_called_once_with("oss-user", flag_keys=[notices.FLAG_KEY])
telemetry._evaluate_mock.assert_called_once_with("oss-user")
telemetry.capture_event.assert_called_once()
event_name, props = telemetry.capture_event.call_args.args
assert event_name == notices.NOTICE_EVENT
@@ -994,7 +997,7 @@ def test_decay_usage_telemetry_disabled_does_not_touch_posthog_or_state(monkeypa
def test_decay_usage_posthog_failure_does_not_consume_cap(notice_harness, capsys):
config, telemetry = notice_harness
telemetry.posthog.evaluate_flags.side_effect = RuntimeError("posthog down")
telemetry._evaluate_mock.side_effect = RuntimeError("posthog down")
notices.display_decay_usage_notice(
MagicMock(),
@@ -1034,7 +1037,7 @@ def test_decay_usage_cap_blocks_before_posthog_eval(notice_harness, capsys):
)
assert capsys.readouterr().err == "Decay usage CTA\n" * notices.DECAY_USAGE_CAP
assert telemetry.posthog.evaluate_flags.call_count == notices.DECAY_USAGE_CAP
assert telemetry._evaluate_mock.call_count == notices.DECAY_USAGE_CAP
assert telemetry.capture_event.call_count == notices.DECAY_USAGE_CAP
assert len(config["notice_state"]["decay_usage"]["events"]) == notices.DECAY_USAGE_CAP
@@ -1111,7 +1114,7 @@ def test_scale_threshold_displayed_logs_and_captures_event(notice_harness, capsy
)
assert capsys.readouterr().err == "Scale top 50\n"
telemetry.posthog.evaluate_flags.assert_called_once_with("oss-user", flag_keys=[notices.FLAG_KEY])
telemetry._evaluate_mock.assert_called_once_with("oss-user")
telemetry.capture_event.assert_called_once()
event_name, props = telemetry.capture_event.call_args.args
assert event_name == notices.NOTICE_EVENT
@@ -1238,7 +1241,7 @@ def test_scale_threshold_telemetry_disabled_does_not_touch_posthog_or_state(monk
def test_scale_threshold_posthog_failure_does_not_consume_cap(notice_harness, capsys):
config, telemetry = notice_harness
telemetry.posthog.evaluate_flags.side_effect = RuntimeError("network unavailable")
telemetry._evaluate_mock.side_effect = RuntimeError("network unavailable")
notices.display_scale_threshold_notice(
MagicMock(),
@@ -1281,7 +1284,7 @@ def test_scale_threshold_cap_blocks_before_posthog_eval(notice_harness, capsys):
)
assert capsys.readouterr().err == "Scale top 50\n" * notices.SCALE_THRESHOLD_CAP
assert telemetry.posthog.evaluate_flags.call_count == notices.SCALE_THRESHOLD_CAP
assert telemetry._evaluate_mock.call_count == notices.SCALE_THRESHOLD_CAP
assert telemetry.capture_event.call_count == notices.SCALE_THRESHOLD_CAP
assert len(config["notice_state"]["scale_threshold"]["events"]) == notices.SCALE_THRESHOLD_CAP
@@ -1549,3 +1552,92 @@ def test_notice_priority_scale_beats_first_run(monkeypatch):
)
assert calls == ["scale"]
class TestStaticFlagEvaluation:
def test_static_flag_result_exposes_posthog_event_properties(self):
flags = notices.StaticFlagResult(notices.DISPLAYED_VARIANT, {"notices": {}})
assert flags._get_event_properties() == {
f"$feature/{notices.FLAG_KEY}": notices.DISPLAYED_VARIANT,
"$active_feature_flags": [notices.FLAG_KEY],
}
def test_evaluate_notice_flags_returns_displayed_or_holdout(self):
result = notices._evaluate_notice_flags("user-abc")
variant = result.get_flag(notices.FLAG_KEY)
assert variant in (notices.DISPLAYED_VARIANT, notices.HOLDOUT_VARIANT)
def test_evaluate_notice_flags_deterministic(self):
a = notices._evaluate_notice_flags("user-123").get_flag(notices.FLAG_KEY)
b = notices._evaluate_notice_flags("user-123").get_flag(notices.FLAG_KEY)
assert a == b
@pytest.mark.parametrize(
("user_id", "expected_variant"),
[("user-0", notices.HOLDOUT_VARIANT), ("user-1", notices.DISPLAYED_VARIANT)],
)
def test_evaluate_notice_flags_preserves_posthog_cohorts(self, monkeypatch, user_id, expected_variant):
monkeypatch.setattr(notices, "_cached_config", {"variant_split": 0.5, "notices": {}})
monkeypatch.setattr(notices, "_cached_config_ts", float("inf"))
assert notices._evaluate_notice_flags(user_id).get_flag(notices.FLAG_KEY) == expected_variant
def test_evaluate_notice_flags_returns_payload_with_notices(self):
result = notices._evaluate_notice_flags("user-abc")
payload = result.get_flag_payload(notices.FLAG_KEY)
assert "notices" in payload
def test_get_notice_config_uses_bundled_fallback(self, monkeypatch):
monkeypatch.setattr(notices, "_cached_config", None)
monkeypatch.setattr(notices, "_cached_config_ts", 0.0)
monkeypatch.setattr(notices, "_fetch_remote_config", lambda: None)
config = notices._get_notice_config()
assert "notices" in config
assert config.get("version") == 1
def test_remote_config_fetch_uses_telemetry_timeout(self, monkeypatch):
timeouts = []
def fail_fetch(_request, timeout):
timeouts.append(timeout)
raise TimeoutError
monkeypatch.setattr(notices.urllib.request, "urlopen", fail_fetch)
assert notices._fetch_remote_config() is None
assert timeouts == [telemetry_module.FEATURE_FLAGS_REQUEST_TIMEOUT_SECONDS]
def test_get_notice_config_prefers_remote(self, monkeypatch):
remote = {"version": 2, "variant_split": 0.5, "notices": {"first_run": {"copy": "remote", "enabled": True}}}
monkeypatch.setattr(notices, "_cached_config", None)
monkeypatch.setattr(notices, "_cached_config_ts", 0.0)
monkeypatch.setattr(notices, "_fetch_remote_config", lambda: remote)
config = notices._get_notice_config()
assert config["version"] == 2
def test_get_notice_config_caches_result(self, monkeypatch):
call_count = 0
remote = {"version": 3, "variant_split": 0.5, "notices": {}}
def counting_fetch():
nonlocal call_count
call_count += 1
return remote
monkeypatch.setattr(notices, "_cached_config", None)
monkeypatch.setattr(notices, "_cached_config_ts", 0.0)
monkeypatch.setattr(notices, "_fetch_remote_config", counting_fetch)
notices._get_notice_config()
notices._get_notice_config()
assert call_count == 1
def test_variant_split_respects_config(self, monkeypatch):
monkeypatch.setattr(notices, "_cached_config", {"variant_split": 1.0, "notices": {}})
monkeypatch.setattr(notices, "_cached_config_ts", float("inf"))
result = notices._evaluate_notice_flags("any-user")
assert result.get_flag(notices.FLAG_KEY) == notices.DISPLAYED_VARIANT
monkeypatch.setattr(notices, "_cached_config", {"variant_split": 0.0, "notices": {}})
result = notices._evaluate_notice_flags("any-user")
assert result.get_flag(notices.FLAG_KEY) == notices.HOLDOUT_VARIANT
@@ -3,8 +3,8 @@ from unittest.mock import AsyncMock, MagicMock
import pytest
from mem0.memory import notices
from mem0.memory import main as memory_main
from mem0.memory import notices
from mem0.memory.main import AsyncMemory, Memory
@@ -34,6 +34,7 @@ def notice_harness(monkeypatch):
config = {}
telemetry = MagicMock()
telemetry.user_id = "oss-user"
evaluate_mock = MagicMock()
def write_config(updated):
saved = deepcopy(updated)
@@ -44,12 +45,14 @@ def notice_harness(monkeypatch):
monkeypatch.setattr(notices, "_write_config", write_config)
monkeypatch.setattr(notices.telemetry_module, "MEM0_TELEMETRY", True)
monkeypatch.setattr(notices.telemetry_module, "_get_oss_telemetry", lambda: telemetry)
monkeypatch.setattr(notices, "_evaluate_notice_flags", evaluate_mock)
telemetry._evaluate_mock = evaluate_mock
return config, telemetry
def configure_flag(telemetry, variant, payload):
flags = FakeFlags(variant, payload)
telemetry.posthog.evaluate_flags.return_value = flags
telemetry._evaluate_mock.return_value = flags
return flags
@@ -274,7 +277,7 @@ def test_performance_slow_query_displayed_logs_and_captures_event(notice_harness
)
assert capsys.readouterr().err == "Performance CTA\n"
telemetry.posthog.evaluate_flags.assert_called_once_with("oss-user", flag_keys=[notices.FLAG_KEY])
telemetry._evaluate_mock.assert_called_once_with("oss-user")
telemetry.capture_event.assert_called_once()
event_name, props = telemetry.capture_event.call_args.args
assert event_name == notices.NOTICE_EVENT
@@ -421,7 +424,7 @@ def test_performance_slow_query_cap_blocks_before_posthog_eval(notice_harness, c
)
assert capsys.readouterr().err == "Performance CTA\n" * notices.PERFORMANCE_SLOW_QUERY_CAP
assert telemetry.posthog.evaluate_flags.call_count == notices.PERFORMANCE_SLOW_QUERY_CAP
assert telemetry._evaluate_mock.call_count == notices.PERFORMANCE_SLOW_QUERY_CAP
assert telemetry.capture_event.call_count == notices.PERFORMANCE_SLOW_QUERY_CAP
assert len(config["notice_state"]["performance_slow_query"]["events"]) == notices.PERFORMANCE_SLOW_QUERY_CAP