From 530535c6ea8f5ee99e2c135afd74fedda05c53b4 Mon Sep 17 00:00:00 2001 From: Adam <2363879+adamdotdevin@users.noreply.github.com> Date: Wed, 26 Aug 2026 10:38:12 -0500 Subject: [PATCH] fix(stats): reduce retention query scan --- .../stats/core/src/domain/inference.test.ts | 14 ++++++----- packages/stats/core/src/domain/inference.ts | 25 ++++++++----------- 2 files changed, 19 insertions(+), 20 deletions(-) diff --git a/packages/stats/core/src/domain/inference.test.ts b/packages/stats/core/src/domain/inference.test.ts index 8eaf1f91896..c2889ff4c66 100644 --- a/packages/stats/core/src/domain/inference.test.ts +++ b/packages/stats/core/src/domain/inference.test.ts @@ -174,14 +174,16 @@ describe("inference stat normalization", () => { expect(queries[0]?.cohortDates).toEqual(["2026-08-10", "2026-08-17"]) expect(queries[0]?.query).toContain("AND product = 'go'") expect(queries[0]?.query).toContain("COUNT(*) AS model_requests") + expect(queries[0]?.query).toContain("SUM(model_requests) AS total_requests") + expect(queries[0]?.query).toContain("MAX(model_requests) AS max_model_requests") + expect(queries[0]?.query).toContain("GROUP BY cohort_date, user_key") + expect(queries[0]?.query).toContain("INNER JOIN user_totals") + expect(queries[0]?.query).toContain("model_usage.model_requests = user_totals.max_model_requests") + expect(queries[0]?.query).toContain("user_totals.total_requests >= 10") expect(queries[0]?.query).toContain( - "SUM(model_requests) OVER (PARTITION BY cohort_date, user_key) AS total_requests", + "CAST(model_usage.model_requests AS double) / NULLIF(user_totals.total_requests, 0) >= 0.8", ) - expect(queries[0]?.query).toContain("ROW_NUMBER() OVER") - expect(queries[0]?.query).toContain("PARTITION BY cohort_date, user_key") - expect(queries[0]?.query).toContain("ORDER BY model_requests DESC, model ASC") - expect(queries[0]?.query).toContain("total_requests >= 10") - expect(queries[0]?.query).toContain("CAST(model_requests AS double) / NULLIF(total_requests, 0) >= 0.8") + expect(queries[0]?.query).not.toContain(" OVER (") expect(queries[0]?.query).toContain("WHEN '2026-08-17' THEN '2026-08-10'") expect(queries[0]?.query).toContain("WHEN '2026-08-24' THEN '2026-08-17'") expect(queries[0]?.query).toContain("started_at >= '2026-08-10T00:00:00.000Z'") diff --git a/packages/stats/core/src/domain/inference.ts b/packages/stats/core/src/domain/inference.ts index bf770844462..a1d1a01625f 100644 --- a/packages/stats/core/src/domain/inference.ts +++ b/packages/stats/core/src/domain/inference.ts @@ -141,25 +141,22 @@ WITH normalized AS ( FROM filtered WHERE activity_week IN (${cohortDates}) GROUP BY activity_week, user_key, provider, model -), ranked_models AS ( +), user_totals AS ( SELECT cohort_date, user_key, - provider, - model, - model_requests, - SUM(model_requests) OVER (PARTITION BY cohort_date, user_key) AS total_requests, - ROW_NUMBER() OVER ( - PARTITION BY cohort_date, user_key - ORDER BY model_requests DESC, model ASC - ) AS model_rank + SUM(model_requests) AS total_requests, + MAX(model_requests) AS max_model_requests FROM model_usage + GROUP BY cohort_date, user_key ), primary_models AS ( - SELECT cohort_date, user_key, provider, model - FROM ranked_models - WHERE model_rank = 1 - AND total_requests >= 10 - AND CAST(model_requests AS double) / NULLIF(total_requests, 0) >= 0.8 + SELECT model_usage.cohort_date, model_usage.user_key, model_usage.provider, model_usage.model + FROM model_usage + INNER JOIN user_totals ON model_usage.cohort_date = user_totals.cohort_date + AND model_usage.user_key = user_totals.user_key + AND model_usage.model_requests = user_totals.max_model_requests + WHERE user_totals.total_requests >= 10 + AND CAST(model_usage.model_requests AS double) / NULLIF(user_totals.total_requests, 0) >= 0.8 ), returned AS ( SELECT ${returnCohortSql} AS cohort_date,