Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions src/interface/storage/Storage.ts
Original file line number Diff line number Diff line change
Expand Up @@ -74,8 +74,8 @@ export interface StorageAdapter {
userID: UserId,
event_type: EventKind,
beforeTimestamp: DateTime,
mode: "production" | "test",
auth: AuthContext,
txn?: unknown
): Promise<number>;
query(request: QueryRequest): Promise<QueryResponse>;
query(request: QueryRequest, auth: AuthContext): Promise<QueryResponse>;
}
6 changes: 3 additions & 3 deletions src/routes/gRPC/payment/createCheckoutLink.ts
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,7 @@ export async function createCheckoutLink(
const custom_price = await calculatePrice(
validatedData.userId,
beforeTimestamp,
mode
auth
);
wideEventBuilder?.setPaymentContext({ priceAmount: custom_price });

Expand Down Expand Up @@ -115,9 +115,9 @@ function validateRequest(
async function calculatePrice(
userId: UserId,
beforeTimestamp: DateTime,
mode: "production" | "test"
auth: AuthContext
): Promise<number> {
const price = await calculatePaymentPrice(userId, beforeTimestamp, mode);
const price = await calculatePaymentPrice(userId, beforeTimestamp, auth);

if (typeof price !== "number" || isNaN(price) || price < 0) {
throw PaymentError.priceCalculationFailed(
Expand Down
10 changes: 8 additions & 2 deletions src/routes/gRPC/query/queryEvents.ts
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import {
AggregationRow,
} from "../../../gen/query/v1/query";
import { queryEventsSchema } from "../../../zod/query";
import { AuthError } from "../../../errors/auth";
import { EventError } from "../../../errors/event";
import { formatZodError } from "../../../utils/formatZodError";
import { StorageAdapterFactory } from "../../../factory";
Expand All @@ -14,7 +15,7 @@ import type {
QueryResponse,
} from "../../../interface/storage/Storage";
import type { WideEventBuilder } from "../../../context/requestContext";
import { apiKeyContextKey } from "../../../context/auth";
import { apiKeyContextKey, type AuthContext } from "../../../context/auth";
import { wideEventContextKey } from "../../../context/requestContext";
import type { ContextUnaryCall } from "../../../interface/types/context.js";

Expand All @@ -33,9 +34,14 @@ export async function queryEvents(
queryConditions: countConditions(queryRequest.where),
});

const auth = call[apiKeyContextKey];
if (!auth) {
return callback?.(AuthError.invalidAPIKey("API key context not found"));
}

const adapter =
await StorageAdapterFactory.getEventStorageAdapter("BASIC_USAGE");
const result = await adapter.query(queryRequest);
const result = await adapter.query(queryRequest, auth);

const response = buildProtoResponse(result, queryRequest);
callback?.(null, response);
Expand Down
7 changes: 4 additions & 3 deletions src/services/pricingService.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,13 +2,14 @@ import { StorageAdapterFactory } from "../factory/EventStorageAdapterFactory";
import { StorageError } from "../errors/storage";
import type { UserId } from "../config/identifiers";
import type { DateTime } from "luxon";
import type { AuthContext } from "../context/auth";
import { getPostgresDB } from "../storage/db/postgres/db";
import { executeInTransaction } from "../storage/adapter/postgres/handlers/addEventUtils";

export async function calculatePaymentPrice(
userId: UserId,
beforeTimestamp: DateTime,
mode: "production" | "test"
auth: AuthContext
): Promise<number> {
const beforeTimestampUtc = beforeTimestamp.toUTC();

Expand All @@ -26,8 +27,8 @@ export async function calculatePaymentPrice(
"calculating payment price",
async (txn) => {
const [sdkPrice, aiPrice] = await Promise.all([
sdkAdapter.price(userId, "BASIC_USAGE", beforeTimestampUtc, mode, txn),
aiAdapter.price(userId, "AI_TOKEN_USAGE", beforeTimestampUtc, mode, txn),
sdkAdapter.price(userId, "BASIC_USAGE", beforeTimestampUtc, auth, txn),
aiAdapter.price(userId, "AI_TOKEN_USAGE", beforeTimestampUtc, auth, txn),
]);

if (typeof sdkPrice !== "number" || isNaN(sdkPrice)) {
Expand Down
10 changes: 5 additions & 5 deletions src/storage/adapter/clickhouse/ClickHouseAdapter.ts
Original file line number Diff line number Diff line change
Expand Up @@ -76,16 +76,16 @@ export class ClickHouseAdapter implements StorageAdapter {
userID: UserId,
event_type: EventKind,
beforeTimestamp: DateTime,
mode: "production" | "test",
auth: AuthContext,
_txn?: unknown
): Promise<number> {
switch (event_type) {
case "BASIC_USAGE": {
return await handlePriceRequestBasicUsage(userID, beforeTimestamp, mode);
return await handlePriceRequestBasicUsage(userID, beforeTimestamp, auth);
}

case "AI_TOKEN_USAGE": {
return await handlePriceRequestAiTokenUsage(userID, beforeTimestamp, mode);
return await handlePriceRequestAiTokenUsage(userID, beforeTimestamp, auth);
}

default: {
Expand All @@ -94,7 +94,7 @@ export class ClickHouseAdapter implements StorageAdapter {
}
}

async query(request: QueryRequest): Promise<QueryResponse> {
return await handleQueryEvents(request);
async query(request: QueryRequest, auth: AuthContext): Promise<QueryResponse> {
return await handleQueryEvents(request, auth);
}
}
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import { DateTime } from "luxon";
import type { UserId } from "../../../../config/identifiers";
import type { AuthContext } from "../../../../context/auth";
import { runClickHousePriceQuery } from "../utils";

const VALUE_EXPR = "JSONExtractInt(metrics, 'debit_amount', 'input') + JSONExtractInt(metrics, 'debit_amount', 'input_cache') + JSONExtractInt(metrics, 'debit_amount', 'output')";
Expand All @@ -9,7 +10,7 @@ const WINDOW_QUERY = `SELECT sum(${VALUE_EXPR}) as total FROM ai_token_usage_eve
export async function handlePriceRequestAiTokenUsage(
userId: UserId,
beforeTimestamp: DateTime,
mode: "production" | "test"
auth: AuthContext
): Promise<number> {
return runClickHousePriceQuery(userId, beforeTimestamp, mode, BASE_QUERY, WINDOW_QUERY, "AI_TOKEN_USAGE");
return runClickHousePriceQuery(userId, beforeTimestamp, auth.mode, BASE_QUERY, WINDOW_QUERY, "AI_TOKEN_USAGE");
}
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import { DateTime } from "luxon";
import type { UserId } from "../../../../config/identifiers";
import type { AuthContext } from "../../../../context/auth";
import { runClickHousePriceQuery } from "../utils";

const BASE_QUERY = "SELECT sum(debit_amount) as total FROM basic_usage_events WHERE user_id = {userId:String} AND mode = {mode:String} AND reported_timestamp < {before:DateTime64(3, 'UTC')}";
Expand All @@ -8,7 +9,7 @@ const WINDOW_QUERY = "SELECT sum(debit_amount) as total FROM basic_usage_events
export async function handlePriceRequestBasicUsage(
userId: UserId,
beforeTimestamp: DateTime,
mode: "production" | "test"
auth: AuthContext
): Promise<number> {
return runClickHousePriceQuery(userId, beforeTimestamp, mode, BASE_QUERY, WINDOW_QUERY, "BASIC_USAGE");
return runClickHousePriceQuery(userId, beforeTimestamp, auth.mode, BASE_QUERY, WINDOW_QUERY, "BASIC_USAGE");
}
4 changes: 3 additions & 1 deletion src/storage/adapter/clickhouse/handlers/queryEvents.ts
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ import type {
QueryResultRow,
QueryFieldName,
} from "../../../../interface/storage/Storage";
import type { AuthContext } from "../../../../context/auth";

interface ChFieldDef {
select: string | null;
Expand Down Expand Up @@ -196,7 +197,8 @@ function buildWhereFromGroup(
}

export async function handleQueryEvents(
request: QueryRequest
request: QueryRequest,
auth: AuthContext
): Promise<QueryResponse> {
const tables = getTablesForRequest(request.where);
if (tables.length === 0) {
Expand Down
5 changes: 3 additions & 2 deletions src/storage/adapter/postgres/handlers/priceRequest.ts
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ import { eq, sum, sql, and, type SQL } from "drizzle-orm";
import type { DateTime } from "luxon";
import type { UserId } from "../../../../config/identifiers";
import type { PgTransaction } from "drizzle-orm/pg-core";
import type { AuthContext } from "../../../../context/auth";

type PriceEventTable =
| typeof basicUsageEventsTable
Expand All @@ -20,7 +21,7 @@ export async function handlePriceRequest(
priceColumn: SQL,
eventType: string,
beforeTimestamp: DateTime,
mode: "production" | "test",
auth: AuthContext,
txn?: PgTransaction<any, any, any>
): Promise<number> {
const db = txn ?? getPostgresDB();
Expand All @@ -36,7 +37,7 @@ export async function handlePriceRequest(

let result;
try {
const baseCondition = sql`${priceTable.reportedTimestamp} > ${usersTable.last_billed_timestamp} AND ${priceTable.userId} = ${userId} AND ${priceTable.mode} = ${mode}`;
const baseCondition = sql`${priceTable.reportedTimestamp} > ${usersTable.last_billed_timestamp} AND ${priceTable.userId} = ${userId} AND ${priceTable.mode} = ${auth.mode}`;
const whereClause = beforeTimestamp
? and(
baseCondition,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,11 +4,12 @@ import { sql } from "drizzle-orm";
import type { DateTime } from "luxon";
import type { UserId } from "../../../../config/identifiers";
import type { PgTransaction } from "drizzle-orm/pg-core";
import type { AuthContext } from "../../../../context/auth";

export async function handlePriceRequestAiTokenUsage(
userId: UserId,
beforeTimestamp: DateTime,
mode: "production" | "test",
auth: AuthContext,
txn?: PgTransaction<any, any, any>
): Promise<number> {
return handlePriceRequest(
Expand All @@ -17,7 +18,7 @@ export async function handlePriceRequestAiTokenUsage(
sql`CAST(${aiTokenUsageEventsTable.metrics}->'debit_amount'->>'input' AS integer) + CAST(${aiTokenUsageEventsTable.metrics}->'debit_amount'->>'input_cache' AS integer) + CAST(${aiTokenUsageEventsTable.metrics}->'debit_amount'->>'output' AS integer)`,
"REQUEST_AI_TOKEN_USAGE",
beforeTimestamp,
mode,
auth,
txn
);
}
Original file line number Diff line number Diff line change
Expand Up @@ -4,11 +4,12 @@ import { sql } from "drizzle-orm";
import type { DateTime } from "luxon";
import type { UserId } from "../../../../config/identifiers";
import type { PgTransaction } from "drizzle-orm/pg-core";
import type { AuthContext } from "../../../../context/auth";

export async function handlePriceRequestBasicUsage(
userId: UserId,
beforeTimestamp: DateTime,
mode: "production" | "test",
auth: AuthContext,
txn?: PgTransaction<any, any, any>
): Promise<number> {
return handlePriceRequest(
Expand All @@ -17,7 +18,7 @@ export async function handlePriceRequestBasicUsage(
sql`${basicUsageEventsTable.debitAmount}`,
"REQUEST_BASIC_USAGE",
beforeTimestamp,
mode,
auth,
txn
);
}
4 changes: 3 additions & 1 deletion src/storage/adapter/postgres/handlers/queryEvents.ts
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@ import type {
QueryResponse,
QueryResultRow,
} from "../../../../interface/storage/Storage";
import type { AuthContext } from "../../../../context/auth";

interface PGFieldDef {
select: string | null;
Expand Down Expand Up @@ -186,7 +187,8 @@ function buildSelectColumns(table: EventTableName): SQL {
}

export async function handleQueryEvents(
request: QueryRequest
request: QueryRequest,
auth: AuthContext
): Promise<QueryResponse> {
const tables = getTablesForRequest(request.where);
if (tables.length === 0) {
Expand Down
10 changes: 5 additions & 5 deletions src/storage/adapter/postgres/postgres.ts
Original file line number Diff line number Diff line change
Expand Up @@ -81,20 +81,20 @@ export class PostgresAdapter implements StorageAdapter {
userID: UserId,
event_type: EventKind,
beforeTimestamp: DateTime,
mode: "production" | "test",
auth: AuthContext,
txn?: unknown
): Promise<number> {
const tx = txn as PgTransaction<any, any, any> | undefined;
switch (event_type) {
case "BASIC_USAGE": {
return await handlePriceRequestBasicUsage(userID, beforeTimestamp, mode, tx);
return await handlePriceRequestBasicUsage(userID, beforeTimestamp, auth, tx);
}

case "AI_TOKEN_USAGE": {
return await handlePriceRequestAiTokenUsage(
userID,
beforeTimestamp,
mode,
auth,
tx
);
}
Expand All @@ -105,7 +105,7 @@ export class PostgresAdapter implements StorageAdapter {
}
}

async query(request: QueryRequest): Promise<QueryResponse> {
return await handleQueryEvents(request);
async query(request: QueryRequest, auth: AuthContext): Promise<QueryResponse> {
return await handleQueryEvents(request, auth);
}
}
Loading