fix(stats): reduce retention query scan

This commit is contained in:
Adam 2026-08-26 10:38:12 -05:00
parent 023620b57e
commit 530535c6ea
No known key found for this signature in database
GPG key ID: 9CB48779AF150E75
2 changed files with 19 additions and 20 deletions

View file

@ -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'")

View file

@ -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,