mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-07-22 12:22:20 +00:00
Compare commits
448 Commits
improved-d
...
fix-compos
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ee79a305ba | ||
|
|
ed2c8c9fe2 | ||
|
|
9161256d31 | ||
|
|
1a8f407142 | ||
|
|
eddc070628 | ||
|
|
f7bd250c97 | ||
|
|
29cfeaed9a | ||
|
|
e2b46dab9f | ||
|
|
632244e54f | ||
|
|
3251664513 | ||
|
|
ee16cd0495 | ||
|
|
70ef3c357c | ||
|
|
e4165f1dab | ||
|
|
41edae52b0 | ||
|
|
cb16da5543 | ||
|
|
d52b727bce | ||
|
|
6a2abd3439 | ||
|
|
1e4ed75179 | ||
|
|
efb5719679 | ||
|
|
3a1902b2b5 | ||
|
|
ba1978f830 | ||
|
|
2f51d62a46 | ||
|
|
1ca344ba7b | ||
|
|
dda9b848ba | ||
|
|
033d8c5a38 | ||
|
|
e73cdffdc5 | ||
|
|
e32b332eb5 | ||
|
|
89b84392a4 | ||
|
|
f37ff30598 | ||
|
|
547f45c185 | ||
|
|
a6e81a83cb | ||
|
|
e2143aa173 | ||
|
|
2eb257c2a6 | ||
|
|
490687bb71 | ||
|
|
966613847e | ||
|
|
d030d86f9a | ||
|
|
d9ab46b3bb | ||
|
|
6e596e3860 | ||
|
|
3c13be20cb | ||
|
|
9d5e14903b | ||
|
|
9e33d3b100 | ||
|
|
c88665f6d1 | ||
|
|
a8157b3e2d | ||
|
|
a07e6723d9 | ||
|
|
5bc7b741bc | ||
|
|
1db37cf084 | ||
|
|
6c5b103149 | ||
|
|
27ec052c17 | ||
|
|
00b813f6a8 | ||
|
|
91e5198a94 | ||
|
|
27f1cc3c42 | ||
|
|
a4d048f2b5 | ||
|
|
752d4f3803 | ||
|
|
164ed775c8 | ||
|
|
9d905758ec | ||
|
|
0d1cb66855 | ||
|
|
6e9932e0ac | ||
|
|
59773e972b | ||
|
|
7fd0cdf987 | ||
|
|
b1947d660a | ||
|
|
d45a9edcba | ||
|
|
4dbbb45240 | ||
|
|
a286b5efa0 | ||
|
|
26a0b04aef | ||
|
|
f7ccc25a7f | ||
|
|
8dd0ece501 | ||
|
|
10969e0719 | ||
|
|
b633455071 | ||
|
|
e280d1f8b1 | ||
|
|
81ad9f604b | ||
|
|
2345691176 | ||
|
|
a7615bc827 | ||
|
|
befdc5307e | ||
|
|
145777ffd0 | ||
|
|
37c2bea93d | ||
|
|
985e765285 | ||
|
|
a4b1330627 | ||
|
|
7b69284812 | ||
|
|
7fb1da7b58 | ||
|
|
38f9923469 | ||
|
|
5b986a3e15 | ||
|
|
56d9ff6b3f | ||
|
|
15b31e7c53 | ||
|
|
d0dcc3219d | ||
|
|
52db6fd308 | ||
|
|
ceca0e8efc | ||
|
|
309bc873e6 | ||
|
|
89109dc209 | ||
|
|
234ea19cad | ||
|
|
489b6eba14 | ||
|
|
66421cd17f | ||
|
|
cd7de3958c | ||
|
|
8c0ac499ef | ||
|
|
1d22155e05 | ||
|
|
ab2fb2afc5 | ||
|
|
884bae9fc4 | ||
|
|
92d581573d | ||
|
|
85a3d3adc0 | ||
|
|
2d7f03b2ed | ||
|
|
37bce70f76 | ||
|
|
038bc14757 | ||
|
|
287cbac5f9 | ||
|
|
28d91227af | ||
|
|
846d894a13 | ||
|
|
a2db3e2d57 | ||
|
|
e43ceb2e43 | ||
|
|
689a07f562 | ||
|
|
da487a850e | ||
|
|
fa6d3c76d0 | ||
|
|
ede1804d4b | ||
|
|
3392e8d4cb | ||
|
|
b5174d9753 | ||
|
|
1f2ff8a99c | ||
|
|
0c60644ba2 | ||
|
|
aca8d43a61 | ||
|
|
7fe4c1963b | ||
|
|
9a0919f149 | ||
|
|
d69ab913d4 | ||
|
|
42b8c332df | ||
|
|
aa682bf8ec | ||
|
|
c53e72e80a | ||
|
|
afd81aeca2 | ||
|
|
4e03145323 | ||
|
|
169686681f | ||
|
|
fa7d2804bb | ||
|
|
ccabc5d06d | ||
|
|
cb68227c88 | ||
|
|
c72fc7dd56 | ||
|
|
8c6d1f89dc | ||
|
|
1ab7e54bd7 | ||
|
|
42dceb0cd6 | ||
|
|
05115c3387 | ||
|
|
c0c8cafd00 | ||
|
|
16d6d66d17 | ||
|
|
3681fc8aab | ||
|
|
d95c09e00c | ||
|
|
3cdef5ec17 | ||
|
|
58e1620347 | ||
|
|
42b5a5e6c8 | ||
|
|
d84d249f2c | ||
|
|
8b6393f794 | ||
|
|
81146710ba | ||
|
|
25f39897d7 | ||
|
|
45bdfaee58 | ||
|
|
a1ae6e94e9 | ||
|
|
eebcc67c85 | ||
|
|
3f9e7f7728 | ||
|
|
a5b5549edd | ||
|
|
22eec0162c | ||
|
|
9f55da9bb8 | ||
|
|
16dce9ea81 | ||
|
|
963ee04619 | ||
|
|
b874b1f01c | ||
|
|
b6cca3d3a0 | ||
|
|
86ebc84f4c | ||
|
|
8a74c0543f | ||
|
|
3d16a9b988 | ||
|
|
40f98b99aa | ||
|
|
7fcae5b08d | ||
|
|
f50cb31749 | ||
|
|
85a1ea4b4c | ||
|
|
0f0f8c40bf | ||
|
|
a40d224ee7 | ||
|
|
a637acd8f4 | ||
|
|
58be0c7976 | ||
|
|
d934f3eead | ||
|
|
74b58e5fa3 | ||
|
|
c1497e0cfe | ||
|
|
30db582321 | ||
|
|
a1018776e9 | ||
|
|
5ef499e2ae | ||
|
|
c7c802c610 | ||
|
|
685368bb0a | ||
|
|
55e240d92a | ||
|
|
7da4ad3818 | ||
|
|
453337cb2c | ||
|
|
9e1934bfda | ||
|
|
31a9730cd8 | ||
|
|
d686e0e851 | ||
|
|
7708ed1c8b | ||
|
|
22ede7636a | ||
|
|
c9a5af0ffa | ||
|
|
072079b33d | ||
|
|
2abc28d295 | ||
|
|
7879b68cad | ||
|
|
9eff340c57 | ||
|
|
933ba105d8 | ||
|
|
205dc6c3b3 | ||
|
|
c1ada39278 | ||
|
|
200b32f1ff | ||
|
|
2b623a86c3 | ||
|
|
a68998e8ff | ||
|
|
905ae67015 | ||
|
|
6fde846be2 | ||
|
|
a833cf429e | ||
|
|
c5a205cf98 | ||
|
|
62185fbd38 | ||
|
|
236854bfe4 | ||
|
|
a7886c528f | ||
|
|
bb97e8dedb | ||
|
|
a7455a48d2 | ||
|
|
a090632804 | ||
|
|
d6b755ff81 | ||
|
|
a1c2a785a1 | ||
|
|
dc8ebe66e0 | ||
|
|
750193e99d | ||
|
|
4ee654331f | ||
|
|
b8fcf5bcc8 | ||
|
|
2ce8981f0c | ||
|
|
b76ff62f1c | ||
|
|
73ebfe8975 | ||
|
|
3dae03d731 | ||
|
|
a7b815b29f | ||
|
|
8a1f909083 | ||
|
|
a6c6d3e034 | ||
|
|
f2a0473fb3 | ||
|
|
33d2ef9ef1 | ||
|
|
f1c3b515fe | ||
|
|
74ddc48ff5 | ||
|
|
6db11ed88f | ||
|
|
4d302baeb5 | ||
|
|
063df51ced | ||
|
|
e97fa60b27 | ||
|
|
42b3840df6 | ||
|
|
a0b2b9466c | ||
|
|
9a553de011 | ||
|
|
901a6a9ba2 | ||
|
|
691927a996 | ||
|
|
8fcecf2c1f | ||
|
|
aed967dc44 | ||
|
|
ed8102d533 | ||
|
|
8fa8475cca | ||
|
|
4723b9db4d | ||
|
|
358ff25899 | ||
|
|
7267bb87b9 | ||
|
|
173f5fbcbd | ||
|
|
9006709f8d | ||
|
|
7d48b36be8 | ||
|
|
8b3fdaa545 | ||
|
|
b43d57df23 | ||
|
|
b9a5a9276e | ||
|
|
a79fdf7212 | ||
|
|
9560050946 | ||
|
|
dd88e9b172 | ||
|
|
e31b45fa9e | ||
|
|
1ed8b29d64 | ||
|
|
59f8d31719 | ||
|
|
93a368b1a2 | ||
|
|
d3dd346853 | ||
|
|
0198569a9a | ||
|
|
7e648cb5c2 | ||
|
|
deb75624f3 | ||
|
|
c5cb562165 | ||
|
|
8a89a38864 | ||
|
|
09e7f1f0bf | ||
|
|
d9d082ad5c | ||
|
|
9cd4ff5c21 | ||
|
|
11eb20a2d1 | ||
|
|
a63e81db06 | ||
|
|
8fc1b6484c | ||
|
|
9bc3feff62 | ||
|
|
cb22968ff3 | ||
|
|
e3bca39815 | ||
|
|
958f28fd82 | ||
|
|
c8c30d7cfd | ||
|
|
2c2124952f | ||
|
|
f56ba92ae8 | ||
|
|
1fa71ee806 | ||
|
|
97d164e181 | ||
|
|
9f3c915dec | ||
|
|
4ff6218140 | ||
|
|
67ea6aec43 | ||
|
|
41f346417a | ||
|
|
8a373276ce | ||
|
|
4857d741de | ||
|
|
eabe17cc26 | ||
|
|
8f298ae9c1 | ||
|
|
e5f3b7d755 | ||
|
|
3796e29cb4 | ||
|
|
812f50ea28 | ||
|
|
ce5fa136c2 | ||
|
|
4643aefb22 | ||
|
|
5aa902fd4a | ||
|
|
db144e0903 | ||
|
|
e20b20dbca | ||
|
|
aa700443d3 | ||
|
|
7e89b5bc6d | ||
|
|
8dbf036558 | ||
|
|
12f2cfecc9 | ||
|
|
a0f3378cf5 | ||
|
|
20b4c3a641 | ||
|
|
88f7ee734f | ||
|
|
326a0086b7 | ||
|
|
eb46651373 | ||
|
|
226188e22d | ||
|
|
89fc48eea2 | ||
|
|
f66d98e6b5 | ||
|
|
cfea4f9cd1 | ||
|
|
c74a877ab5 | ||
|
|
83a76e86fe | ||
|
|
66975ee271 | ||
|
|
5788d63892 | ||
|
|
d0c7cc6bd9 | ||
|
|
40096867d9 | ||
|
|
e002c0b66f | ||
|
|
608549d051 | ||
|
|
72f389c8d2 | ||
|
|
97daddb95a | ||
|
|
847f4b07a5 | ||
|
|
4bce04ada4 | ||
|
|
f1e2448620 | ||
|
|
28340a152c | ||
|
|
e28f6118e7 | ||
|
|
9fb6f54d12 | ||
|
|
3cb8d7b5dd | ||
|
|
ede076b881 | ||
|
|
449a0951f9 | ||
|
|
8c9ede2272 | ||
|
|
e794b09614 | ||
|
|
24f6519267 | ||
|
|
f57beb6411 | ||
|
|
3e41e59a1d | ||
|
|
002d750830 | ||
|
|
8c2eb55760 | ||
|
|
350714f23a | ||
|
|
9cc2e84f7c | ||
|
|
6fbd479bdc | ||
|
|
f67c26935a | ||
|
|
9ce91f58a9 | ||
|
|
bb82361434 | ||
|
|
5f1d67e87e | ||
|
|
b2106ad1e6 | ||
|
|
709a4ba0dc | ||
|
|
21f421b212 | ||
|
|
b3c5e4cbf6 | ||
|
|
1d95379328 | ||
|
|
342a7f7f16 | ||
|
|
7694f20006 | ||
|
|
3ea0a267dd | ||
|
|
68e537f6fb | ||
|
|
814c39898d | ||
|
|
fdd4f12f5c | ||
|
|
6957d8c0d9 | ||
|
|
2c358276bf | ||
|
|
48529f672a | ||
|
|
030c2f65e4 | ||
|
|
3510402af2 | ||
|
|
5f376d716d | ||
|
|
d889274f84 | ||
|
|
0c0f19d854 | ||
|
|
b61bffc666 | ||
|
|
af658136d4 | ||
|
|
8973627b5f | ||
|
|
25f427033a | ||
|
|
d3dd8318e4 | ||
|
|
c5fd386c1e | ||
|
|
52c7f17215 | ||
|
|
2c03302055 | ||
|
|
c3221f2a31 | ||
|
|
ce9834d7ec | ||
|
|
6ebe73f2f7 | ||
|
|
98aecb08f9 | ||
|
|
58fa063c6b | ||
|
|
7c94f60797 | ||
|
|
4cb4c6dfec | ||
|
|
512b686e5f | ||
|
|
f495a10eeb | ||
|
|
d1692edb63 | ||
|
|
8f81bcd2fc | ||
|
|
6a5ed9d063 | ||
|
|
b9890e6ad5 | ||
|
|
e4b8293d41 | ||
|
|
ba5f9fc181 | ||
|
|
795fff61e0 | ||
|
|
f9bfd4f0d2 | ||
|
|
5683382ada | ||
|
|
c92372dafd | ||
|
|
b58dd78fde | ||
|
|
4c31bf9767 | ||
|
|
42efa3c1ba | ||
|
|
74b1d39d5c | ||
|
|
4df4976f44 | ||
|
|
4aa57959bf | ||
|
|
1af39f043f | ||
|
|
1751cd3b47 | ||
|
|
b1facd58d5 | ||
|
|
6288d6fef7 | ||
|
|
a1223ad610 | ||
|
|
c75f170ed0 | ||
|
|
c6e401c3f6 | ||
|
|
1b3b206a20 | ||
|
|
248937e05f | ||
|
|
b9b477e5eb | ||
|
|
ace8cf960c | ||
|
|
3396e0cd47 | ||
|
|
31898192b2 | ||
|
|
e50facc835 | ||
|
|
55dc485705 | ||
|
|
855d60b4a5 | ||
|
|
27ace348b5 | ||
|
|
80559a57d5 | ||
|
|
8e9f6647e7 | ||
|
|
73e3d34623 | ||
|
|
c8f8857f03 | ||
|
|
c9650441bb | ||
|
|
5ce9c2217f | ||
|
|
89b8488ab8 | ||
|
|
2ee917fa31 | ||
|
|
5e12a7e92d | ||
|
|
1d043cd98d | ||
|
|
39e0959fcd | ||
|
|
29be9d5b9c | ||
|
|
1ebb7d71e1 | ||
|
|
bbf1e65a5d | ||
|
|
51c3e5dcd7 | ||
|
|
788075f656 | ||
|
|
0bbcacd186 | ||
|
|
6d5b811c20 | ||
|
|
42258ae39c | ||
|
|
04d6903369 | ||
|
|
b0b2ceb1a0 | ||
|
|
bf91f401af | ||
|
|
60e0eebd05 | ||
|
|
8fff716bb2 | ||
|
|
b2ef15a406 | ||
|
|
f8fdcf3bf3 | ||
|
|
9dfa58d69f | ||
|
|
22ec9c3132 | ||
|
|
b72c578954 | ||
|
|
7e299180fe | ||
|
|
f290534df4 | ||
|
|
8d3c064b29 | ||
|
|
4ac96ade5f | ||
|
|
180a469399 | ||
|
|
87fbb48ca8 | ||
|
|
db021866d8 | ||
|
|
d4339287be | ||
|
|
88bcc0edcb | ||
|
|
917a4d32b1 | ||
|
|
daf17f51ab | ||
|
|
fc042c768c | ||
|
|
cb36189db3 | ||
|
|
d2487f42b0 | ||
|
|
5255fce7b2 | ||
|
|
ee668ee93b | ||
|
|
cf8b990fc7 | ||
|
|
78dd74845b | ||
|
|
e7f4c98475 |
@@ -14,6 +14,7 @@ UPSTREAM_API_KEY=your-upstream-api-key
|
||||
# HTTP_URL=https://api.mynode.com
|
||||
# ONION_URL=http://mynode.onion (auto fetched from compose)
|
||||
# RELAYS="wss://relay.damus.io,wss://relay.nostr.band,wss://eden.nostr.land,wss://relay.routstr.com"
|
||||
# ENABLE_ANALYTICS_SHARING=true
|
||||
# CASHU_MINTS="https://mint.minibits.cash/Bitcoin,https://mint.cubabitcoin.org,https://ecashmint.otrta.me"
|
||||
# RECEIVE_LN_ADDRESS=
|
||||
|
||||
|
||||
12
.github/workflows/container.yml
vendored
12
.github/workflows/container.yml
vendored
@@ -14,6 +14,14 @@ jobs:
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v3
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Resolve git metadata
|
||||
id: gitmeta
|
||||
run: |
|
||||
echo "sha=$(git rev-parse --short=7 HEAD)" >> "$GITHUB_OUTPUT"
|
||||
echo "tag=$(git describe --tags --exact-match HEAD 2>/dev/null || true)" >> "$GITHUB_OUTPUT"
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v2
|
||||
@@ -29,7 +37,11 @@ jobs:
|
||||
uses: docker/build-push-action@v4
|
||||
with:
|
||||
context: .
|
||||
file: Dockerfile.full
|
||||
push: true
|
||||
build-args: |
|
||||
GIT_COMMIT=${{ steps.gitmeta.outputs.sha }}
|
||||
GIT_TAG=${{ steps.gitmeta.outputs.tag }}
|
||||
tags: |
|
||||
ghcr.io/routstr/proxy:latest
|
||||
ghcr.io/routstr/core:latest
|
||||
|
||||
2
.github/workflows/test.yml
vendored
2
.github/workflows/test.yml
vendored
@@ -67,7 +67,7 @@ jobs:
|
||||
- name: Setup Node.js
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: "18"
|
||||
node-version: "20"
|
||||
cache: "pnpm"
|
||||
cache-dependency-path: ui/pnpm-lock.yaml
|
||||
|
||||
|
||||
28
Dockerfile
28
Dockerfile
@@ -1,26 +1,28 @@
|
||||
FROM ghcr.io/astral-sh/uv:python3.11-alpine
|
||||
FROM ghcr.io/astral-sh/uv:python3.11-bookworm-slim
|
||||
|
||||
# Install system dependencies required for secp256k1
|
||||
RUN apk add --no-cache \
|
||||
pkgconf \
|
||||
build-base \
|
||||
automake \
|
||||
autoconf \
|
||||
libtool \
|
||||
m4 \
|
||||
perl
|
||||
RUN apk add git
|
||||
RUN apt-get update \
|
||||
&& apt-get install -y --no-install-recommends \
|
||||
build-essential \
|
||||
pkg-config \
|
||||
libsecp256k1-dev \
|
||||
autoconf \
|
||||
automake \
|
||||
libtool \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
COPY uv.lock pyproject.toml ./
|
||||
RUN mkdir -p /routstr
|
||||
|
||||
RUN uv add git+https://github.com/saschanaz/secp256k1-py.git#branch=upgrade060
|
||||
# RUN uv sync
|
||||
RUN uv sync --frozen --no-dev --no-install-project
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
COPY . .
|
||||
|
||||
ARG GIT_COMMIT=""
|
||||
ARG GIT_TAG=""
|
||||
ENV GIT_COMMIT=${GIT_COMMIT}
|
||||
ENV GIT_TAG=${GIT_TAG}
|
||||
ENV PORT=8000
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
|
||||
|
||||
52
Dockerfile.full
Normal file
52
Dockerfile.full
Normal file
@@ -0,0 +1,52 @@
|
||||
# Multi-stage Dockerfile for Routstr (includes UI build)
|
||||
# Stage 1: Build the UI
|
||||
FROM node:23-alpine AS ui-builder
|
||||
WORKDIR /app/ui
|
||||
|
||||
# Install pnpm
|
||||
RUN corepack enable pnpm && corepack prepare pnpm@10.15.0 --activate
|
||||
|
||||
# Copy UI source
|
||||
COPY ui/package.json ui/pnpm-lock.yaml* ./
|
||||
RUN pnpm install --frozen-lockfile
|
||||
|
||||
COPY ui/ ./
|
||||
ENV NEXT_TELEMETRY_DISABLED=1
|
||||
# Next.js build produces a static export in 'out' directory
|
||||
RUN pnpm run build
|
||||
|
||||
# Stage 2: Build the Routstr Node
|
||||
FROM ghcr.io/astral-sh/uv:python3.11-bookworm-slim AS runner
|
||||
|
||||
RUN apt-get update \
|
||||
&& apt-get install -y --no-install-recommends \
|
||||
build-essential \
|
||||
pkg-config \
|
||||
libsecp256k1-dev \
|
||||
autoconf \
|
||||
automake \
|
||||
libtool \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
COPY uv.lock pyproject.toml ./
|
||||
|
||||
RUN uv sync --no-dev --no-install-project
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
COPY . .
|
||||
|
||||
# Copy the built UI from the ui-builder stage
|
||||
COPY --from=ui-builder /app/ui/out ./ui_out
|
||||
|
||||
ARG GIT_COMMIT=""
|
||||
ARG GIT_TAG=""
|
||||
ENV GIT_COMMIT=${GIT_COMMIT}
|
||||
ENV GIT_TAG=${GIT_TAG}
|
||||
ENV PORT=8000
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
|
||||
EXPOSE 8000
|
||||
|
||||
# Run the application
|
||||
CMD ["/.venv/bin/fastapi", "run", "routstr", "--host", "0.0.0.0"]
|
||||
36
README.md
36
README.md
@@ -13,7 +13,7 @@ This repo contains Routstr Core: a FastAPI-based reverse proxy that sits in fron
|
||||
|
||||
- **Overview**: <https://docs.routstr.com/overview/>
|
||||
- **Provider Guide**: <https://docs.routstr.com/provider/quickstart/>
|
||||
- **User Guide**: <https://docs.routstr.com/user-guide/>
|
||||
- **User Guide**: <https://docs.routstr.com/user-guide/introduction/>
|
||||
|
||||
## Basic Usage
|
||||
|
||||
@@ -40,7 +40,7 @@ print(response.choices[0].message.content)
|
||||
### cURL
|
||||
|
||||
```bash
|
||||
curl http://localhost:8000/v1/chat/completions \
|
||||
curl https://api.routstr.com/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "x-cashu: cashuBo2FteCJodHRwczovL21..." \
|
||||
-d '{
|
||||
@@ -51,23 +51,31 @@ curl http://localhost:8000/v1/chat/completions \
|
||||
|
||||
## Quick Start (Docker)
|
||||
|
||||
If you are a node runner, start a Routstr Core instance and configure upstream access in the dashboard.
|
||||
If you are a node runner, start a Routstr Core instance using Docker Compose:
|
||||
|
||||
```bash
|
||||
docker run -d \
|
||||
--name routstr-proxy \
|
||||
-p 8000:8000 \
|
||||
ghcr.io/routstr/proxy:latest
|
||||
```
|
||||
1. **Prepare your `.env`**:
|
||||
```bash
|
||||
ADMIN_PASSWORD=mysecretpassword
|
||||
NAME="My AI Node"
|
||||
DESCRIPTION="Fast access to models"
|
||||
NSEC=yournsec
|
||||
RECEIVE_LN_ADDRESS=yourname@wallet.com
|
||||
```
|
||||
|
||||
2. **Start the services**:
|
||||
```bash
|
||||
docker compose up -d
|
||||
```
|
||||
|
||||
3. **Configure**:
|
||||
Open [http://localhost:8000/admin/](http://localhost:8000/admin/) to connect your AI providers and set pricing.
|
||||
|
||||
For full instructions, see the **[Provider Quick Start Guide](https://docs.routstr.com/provider/quickstart/)**.
|
||||
|
||||
## Development
|
||||
|
||||
```bash
|
||||
make setup
|
||||
cp .env.example .env
|
||||
fastapi run routstr --host 0.0.0.0 --port 8000
|
||||
fastapi run routstr
|
||||
```
|
||||
|
||||
## License
|
||||
|
||||
GPLv3. See `LICENSE`.
|
||||
|
||||
11
compose.yml
11
compose.yml
@@ -6,10 +6,11 @@ services:
|
||||
context: ./ui
|
||||
dockerfile: Dockerfile.build
|
||||
args:
|
||||
NEXT_PUBLIC_API_URL: ${NEXT_PUBLIC_API_URL:-http://127.0.0.1:8000}
|
||||
# NEXT_PUBLIC_API_URL: ${NEXT_PUBLIC_API_URL:-http://127.0.0.1:8000}
|
||||
NEXT_PUBLIC_ADMIN_API_KEY: ${NEXT_PUBLIC_ADMIN_API_KEY:-}
|
||||
user: root
|
||||
volumes:
|
||||
- ./ui_out:/output
|
||||
- ./ui_out:/output:z
|
||||
command:
|
||||
["sh", "-c", "mkdir -p /output && cp -r /app/built/. /output/ && echo 'UI build copied to mounted volume' && ls -la /output/ && echo 'UI built and ready' && tail -f /dev/null"]
|
||||
|
||||
@@ -18,10 +19,10 @@ services:
|
||||
depends_on:
|
||||
- ui
|
||||
volumes:
|
||||
- .:/app
|
||||
- ./logs:/app/logs
|
||||
- .:/app:z
|
||||
- ./logs:/app/logs:z
|
||||
- tor-data:/var/lib/tor:ro
|
||||
- ./ui_out:/app/ui_out:ro
|
||||
- ./ui_out:/app/ui_out:ro,z
|
||||
env_file:
|
||||
- .env
|
||||
environment:
|
||||
|
||||
@@ -360,6 +360,44 @@ POST /v1/wallet/create
|
||||
}
|
||||
```
|
||||
|
||||
### Get Key Information
|
||||
|
||||
Get current balance, consumption data, and child keys for an API key.
|
||||
|
||||
```http
|
||||
GET /v1/balance/info
|
||||
Authorization: Bearer sk-...
|
||||
```
|
||||
|
||||
**Response:**
|
||||
|
||||
```json
|
||||
{
|
||||
"api_key": "sk-abc...",
|
||||
"balance": 8500000,
|
||||
"reserved": 0,
|
||||
"is_child": false,
|
||||
"parent_key": null,
|
||||
"total_requests": 42,
|
||||
"total_spent": 1500000,
|
||||
"balance_limit": null,
|
||||
"balance_limit_reset": null,
|
||||
"validity_date": null,
|
||||
"child_keys": [
|
||||
{
|
||||
"api_key": "sk-child1...",
|
||||
"total_requests": 10,
|
||||
"total_spent": 500000,
|
||||
"balance_limit": 1000000,
|
||||
"balance_limit_reset": "daily",
|
||||
"validity_date": 1738000000
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
`balance` is the spendable balance used by request admission.
|
||||
|
||||
### Check Balance
|
||||
|
||||
Get current wallet balance.
|
||||
@@ -434,6 +472,42 @@ Authorization: Bearer sk-...
|
||||
}
|
||||
```
|
||||
|
||||
### Create Child Key
|
||||
|
||||
Creates one or more child API keys that share the parent's balance. Each child key creation costs a fixed amount (configurable).
|
||||
|
||||
```http
|
||||
POST /v1/balance/child-key
|
||||
Authorization: Bearer sk-...
|
||||
```
|
||||
|
||||
**Request Body:**
|
||||
|
||||
```json
|
||||
{
|
||||
"count": 1
|
||||
}
|
||||
```
|
||||
|
||||
**Parameters:**
|
||||
|
||||
| Parameter | Type | Required | Default | Description |
|
||||
|-----------|------|----------|---------|-------------|
|
||||
| `count` | integer | Yes | - | Number of child keys to create (1-50) |
|
||||
|
||||
**Response:**
|
||||
|
||||
```json
|
||||
{
|
||||
"api_keys": ["sk-abc...", "sk-def..."],
|
||||
"count": 2,
|
||||
"cost_msats": 2000,
|
||||
"cost_sats": 2,
|
||||
"parent_balance": 98000,
|
||||
"parent_balance_sats": 98
|
||||
}
|
||||
```
|
||||
|
||||
## Provider Discovery
|
||||
|
||||
## Admin Settings
|
||||
|
||||
@@ -17,7 +17,7 @@ Cashu ([cashu.me](https://cashu.me)) or Lightning ([Strike](https://strike.me),
|
||||
|
||||
### 🌐 Provider
|
||||
|
||||
A Routstr node, e.g. `https://api.routstr.com/v1`
|
||||
A Routstr node, e.g. `https://api.routstr.com`
|
||||
|
||||
### 🤖 Client
|
||||
|
||||
|
||||
@@ -53,9 +53,11 @@ If your balance runs low, you don't need a new key. You can top up the existing
|
||||
|
||||
### Via Lightning
|
||||
|
||||
`POST /lightning/invoice` with `{"amount_sats": 1000, "purpose": "topup", "api_key": "sk-..."}`.
|
||||
`POST /lightning/invoice` with `Authorization: Bearer sk-...` header and body `{"amount_sats": 1000, "purpose": "topup"}`.
|
||||
*Once paid, the funds are added to your existing key.*
|
||||
|
||||
> Legacy: the endpoint is also exposed at `/v1/balance/lightning/invoice`, and accepts an `api_key` field in the body as a fallback for older clients. New integrations should use the RIP-08 path with the `Authorization` header.
|
||||
|
||||
### Via Cashu
|
||||
|
||||
`POST /v1/balance/topup` with `{"cashu_token": "..."}` and `Authorization: Bearer sk-...`.
|
||||
|
||||
@@ -110,5 +110,5 @@ Higher margins, fewer clients:
|
||||
|
||||
### Mixed Strategy
|
||||
|
||||
- Cheap models (GPT-3.5, Haiku): Low margin to attract volume
|
||||
- Premium models (GPT-4, Opus): High margin for profit
|
||||
- Cheap models (GLM-4.7-Flash, Seed-1.6): Low margin to attract volume
|
||||
- Premium models (GPT-5-Pro, Claude-Opus): High margin for profit
|
||||
|
||||
@@ -6,6 +6,32 @@ For automated deployments, you can optionally pre-configure settings via environ
|
||||
|
||||
---
|
||||
|
||||
## Initial Setup (.env file)
|
||||
|
||||
Before running your node, you should create a `.env` file in the project root. This file is used to bootstrap the initial configuration and store sensitive secrets.
|
||||
|
||||
### Example .env
|
||||
|
||||
```bash
|
||||
ADMIN_PASSWORD=your-secure-password
|
||||
|
||||
# Node Identity
|
||||
NAME="My AI Node"
|
||||
DESCRIPTION="Fast access to models"
|
||||
|
||||
# Lightning Payouts
|
||||
RECEIVE_LN_ADDRESS=yourname@wallet.com
|
||||
```
|
||||
|
||||
### Setting the UI Password
|
||||
|
||||
There are two ways to set or change your Admin Dashboard password:
|
||||
|
||||
1. **Via Environment Variable**: Set `ADMIN_PASSWORD` in your `.env` file before starting the container. This will be the password used for the first login.
|
||||
2. **Via Dashboard**: Once logged in, go to **Settings** → **Security** to update your password. Dashboard settings override the `.env` file once saved.
|
||||
|
||||
---
|
||||
|
||||
## Admin Dashboard (Primary)
|
||||
|
||||
Access the dashboard at `/admin/` on your node.
|
||||
@@ -14,29 +40,29 @@ Access the dashboard at `/admin/` on your node.
|
||||
|
||||
Connect to your AI provider(s):
|
||||
|
||||
| Setting | Description |
|
||||
|---------|-------------|
|
||||
| Setting | Description |
|
||||
| ---------------- | ------------------------------------------------ |
|
||||
| **Upstream URL** | API endpoint (e.g., `https://api.openai.com/v1`) |
|
||||
| **API Key** | Your provider's API key |
|
||||
| **API Key** | Your provider's API key |
|
||||
|
||||
### Node Identity
|
||||
|
||||
How your node appears to clients:
|
||||
|
||||
| Setting | Description |
|
||||
|---------|-------------|
|
||||
| **Name** | Display name (e.g., "Fast GPT-4 Node") |
|
||||
| **Description** | Brief description of your service |
|
||||
| Setting | Description |
|
||||
| --------------- | -------------------------------------- |
|
||||
| **Name** | Display name (e.g., "Fast GPT-4 Node") |
|
||||
| **Description** | Brief description of your service |
|
||||
|
||||
### Pricing
|
||||
|
||||
Control your profit margins:
|
||||
|
||||
| Setting | Description | Default |
|
||||
|---------|-------------|---------|
|
||||
| **Fixed Pricing** | Charge flat rate per request vs. per-token | Off |
|
||||
| **Exchange Fee** | Buffer for BTC volatility | 1.005 (0.5%) |
|
||||
| **Upstream Fee** | Your profit markup | 1.10 (10%) |
|
||||
| Setting | Description | Default |
|
||||
| ----------------- | ------------------------------------------ | ------------ |
|
||||
| **Fixed Pricing** | Charge flat rate per request vs. per-token | Off |
|
||||
| **Exchange Fee** | Buffer for BTC volatility | 1.005 (0.5%) |
|
||||
| **Upstream Fee** | Your profit markup | 1.10 (10%) |
|
||||
|
||||
See [Pricing](pricing.md) for detailed strategies.
|
||||
|
||||
@@ -44,33 +70,40 @@ See [Pricing](pricing.md) for detailed strategies.
|
||||
|
||||
Which mints to accept payments from:
|
||||
|
||||
| Setting | Description |
|
||||
|---------|-------------|
|
||||
| Setting | Description |
|
||||
| --------- | ------------------------------- |
|
||||
| **Mints** | List of trusted Cashu mint URLs |
|
||||
|
||||
### Lightning Withdrawals
|
||||
|
||||
Automatic profit withdrawal:
|
||||
|
||||
| Setting | Description |
|
||||
|---------|-------------|
|
||||
| **Lightning Address** | Your LN address for withdrawals |
|
||||
| Setting | Description | Default |
|
||||
| ------------------------------------ | ---------------------------------------------------------------------------------------------------------- | ------- |
|
||||
| **Lightning Address** | Your LN address for withdrawals | — |
|
||||
| **Minimum Payout (sat)** | Min available balance (in sats) before profit is paid out. Applies to both `sat` and `msat` mints (auto-converted). | `210` |
|
||||
| **Payout Interval (seconds)** | How often the payout loop wakes up and checks balances | `900` |
|
||||
|
||||
All payout amounts must be positive. Set the minimums above your wallet's
|
||||
minimum-invoice constraint (typically 1 sat) and high enough to amortise
|
||||
routing fees.
|
||||
|
||||
### Security
|
||||
|
||||
| Setting | Description |
|
||||
|---------|-------------|
|
||||
| Setting | Description |
|
||||
| ------------------ | ----------------------------- |
|
||||
| **Admin Password** | Password for dashboard access |
|
||||
|
||||
### Nostr Discovery
|
||||
|
||||
Announce your node on the network:
|
||||
|
||||
| Setting | Description |
|
||||
|---------|-------------|
|
||||
| **Npub** | Your Nostr public key |
|
||||
| **Nsec** | Your Nostr private key (for signing) |
|
||||
| **Relays** | Relays to publish announcements |
|
||||
| Setting | Description |
|
||||
| ---------- | ------------------------------------ |
|
||||
| **Npub** | Your Nostr public key |
|
||||
| **Nsec** | Your Nostr private key (for signing) |
|
||||
| **Relays** | Relays to publish announcements |
|
||||
| **Share Analytics** | Publish aggregate usage stats to Nostr |
|
||||
|
||||
See [Discovery](discovery.md) for details.
|
||||
|
||||
@@ -86,21 +119,24 @@ Use environment variables for:
|
||||
|
||||
### All Variables
|
||||
|
||||
| Variable | Description | Default |
|
||||
|----------|-------------|---------|
|
||||
| `UPSTREAM_BASE_URL` | Upstream API endpoint | — |
|
||||
| `UPSTREAM_API_KEY` | Upstream API key | — |
|
||||
| `ADMIN_PASSWORD` | Dashboard password | (none) |
|
||||
| `DATABASE_URL` | Database connection string | `sqlite+aiosqlite:///keys.db` |
|
||||
| `NAME` | Node display name | `ARoutstrNode` |
|
||||
| `DESCRIPTION` | Node description | `A Routstr Node` |
|
||||
| `NPUB` | Nostr public key (bech32) | — |
|
||||
| `NSEC` | Nostr private key | — |
|
||||
| `CASHU_MINTS` | Comma-separated mint URLs | `https://mint.minibits.cash/Bitcoin` |
|
||||
| `RECEIVE_LN_ADDRESS` | Lightning address for withdrawals | — |
|
||||
| `TOR_PROXY_URL` | SOCKS5 proxy for Tor | `socks5://127.0.0.1:9050` |
|
||||
| `CORS_ORIGINS` | Allowed CORS origins | `*` |
|
||||
| `RELAYS` | Nostr relays (comma-separated) | (default set) |
|
||||
| Variable | Description | Default |
|
||||
| -------------------- | --------------------------------- | ------------------------------------ |
|
||||
| `UPSTREAM_BASE_URL` | Upstream API endpoint | — |
|
||||
| `UPSTREAM_API_KEY` | Upstream API key | — |
|
||||
| `ADMIN_PASSWORD` | Dashboard password | (none) |
|
||||
| `DATABASE_URL` | Database connection string | `sqlite+aiosqlite:///keys.db` |
|
||||
| `NAME` | Node display name | `ARoutstrNode` |
|
||||
| `DESCRIPTION` | Node description | `A Routstr Node` |
|
||||
| `NPUB` | Nostr public key (bech32) | — |
|
||||
| `NSEC` | Nostr private key | — |
|
||||
| `ENABLE_ANALYTICS_SHARING` | Enable usage analytics sharing to Nostr | `true` |
|
||||
| `CASHU_MINTS` | Comma-separated mint URLs | `https://mint.minibits.cash/Bitcoin` |
|
||||
| `RECEIVE_LN_ADDRESS` | Lightning address for withdrawals | — |
|
||||
| `MIN_PAYOUT_SAT` | Min payout balance in sats (applies to all mints) | `210` |
|
||||
| `PAYOUT_INTERVAL_SECONDS` | Payout loop interval (seconds) | `900` |
|
||||
| `TOR_PROXY_URL` | SOCKS5 proxy for Tor | `socks5://127.0.0.1:9050` |
|
||||
| `CORS_ORIGINS` | Allowed CORS origins | `*` |
|
||||
| `RELAYS` | Nostr relays (comma-separated) | (default set) |
|
||||
|
||||
### Priority
|
||||
|
||||
|
||||
@@ -152,6 +152,7 @@ Manage which mints you accept payments from:
|
||||
|-------|-------------|
|
||||
| **Nsec** | Private key for signing announcements |
|
||||
| **Relays** | Where to publish your node advertisement |
|
||||
| **Share Analytics** | Toggle publishing aggregate usage stats to Nostr |
|
||||
|
||||
### Security
|
||||
|
||||
|
||||
@@ -2,34 +2,71 @@
|
||||
|
||||
Production deployment guide for Routstr Provider nodes.
|
||||
|
||||
## Docker Compose (Recommended)
|
||||
## All-in-One Docker Image (Preferred)
|
||||
|
||||
For production, use Docker Compose with persistent storage and optional Tor support.
|
||||
The easiest way to deploy Routstr is using the all-in-one Docker image from Docker Hub, which includes both the FastAPI backend and the Next.js admin dashboard in a single container.
|
||||
|
||||
### Basic Setup
|
||||
### Quick Start
|
||||
|
||||
Create a `compose.yml`:
|
||||
```bash
|
||||
docker run -d \
|
||||
--name routstr \
|
||||
-p 8000:8000 \
|
||||
-v routstr-data:/app/data \
|
||||
-e DATABASE_URL="sqlite:////app/data/routstr.db" \
|
||||
9qeklajc/routstr:latest
|
||||
```
|
||||
|
||||
Access your node:
|
||||
- **API & Admin Dashboard**: http://localhost:8000
|
||||
|
||||
### Docker Compose Setup
|
||||
|
||||
Create `docker-compose.yml`:
|
||||
|
||||
```yaml
|
||||
version: '3.8'
|
||||
|
||||
services:
|
||||
routstr:
|
||||
image: ghcr.io/routstr/proxy:latest
|
||||
image: 9qeklajc/routstr:latest
|
||||
container_name: routstr
|
||||
restart: unless-stopped
|
||||
ports:
|
||||
- "8000:8000"
|
||||
volumes:
|
||||
- ./data:/app/data
|
||||
- ./logs:/app/logs
|
||||
- routstr-data:/app/data
|
||||
environment:
|
||||
DATABASE_URL: "sqlite:////app/data/routstr.db"
|
||||
ADMIN_KEY: "your-secure-admin-key"
|
||||
LOG_LEVEL: "info"
|
||||
|
||||
volumes:
|
||||
routstr-data:
|
||||
```
|
||||
|
||||
Start the node:
|
||||
Start it:
|
||||
|
||||
```bash
|
||||
docker compose up -d
|
||||
```
|
||||
|
||||
Then configure everything via the [Admin Dashboard](http://localhost:8000/admin/).
|
||||
---
|
||||
|
||||
## Docker Compose (Recommended)
|
||||
|
||||
For production, use Docker Compose with persistent storage and optional Tor support.
|
||||
|
||||
Use the included `compose.yml` for a flexible setup that handles both the UI and the node execution. This is useful for development or when you want to manage Tor as a separate service.
|
||||
|
||||
```bash
|
||||
docker compose up -d
|
||||
```
|
||||
|
||||
This will:
|
||||
1. **Build the UI**: Compiles the frontend and copies it to a shared volume.
|
||||
2. **Start Routstr**: Runs the Python node, mounting the built UI.
|
||||
3. **Start Tor**: Provides anonymous access via a `.onion` address.
|
||||
|
||||
---
|
||||
|
||||
@@ -189,8 +226,16 @@ docker compose up -d
|
||||
|
||||
## Building from Source
|
||||
|
||||
### Using Docker Compose
|
||||
The easiest way to build everything from source:
|
||||
|
||||
```bash
|
||||
git clone https://github.com/routstr/routstr-core.git
|
||||
cd routstr-core
|
||||
docker build -t routstr-local .
|
||||
docker compose build
|
||||
```
|
||||
|
||||
### Individual Components
|
||||
If you prefer building the node only (requires manual UI build first):
|
||||
|
||||
```bash
|
||||
docker build -t routstr-node .
|
||||
```
|
||||
|
||||
@@ -13,7 +13,7 @@ A **Routstr Provider Node** acts as a gateway that:
|
||||
You bring the API keys, Routstr handles the billing, payments, and client management.
|
||||
|
||||
!!! tip "Future: Node-to-Node Routing"
|
||||
In future versions, you'll be able to run a node that connects to other Routstr nodes—eliminating the need to configure upstream providers yourself. For now, you'll need your own API credentials.
|
||||
In future versions, you'll be able to run a node that connects to other Routstr nodes—eliminating the need to configure upstream providers yourself. For now, you'll need your own API credentials.
|
||||
|
||||
---
|
||||
|
||||
@@ -24,14 +24,30 @@ You bring the API keys, Routstr handles the billing, payments, and client manage
|
||||
|
||||
---
|
||||
|
||||
## 1. Start the Node
|
||||
## 1. Prepare Configuration
|
||||
|
||||
Create a `.env` file in the root of the project to store your secrets:
|
||||
|
||||
```bash
|
||||
docker run -d \
|
||||
--name routstr \
|
||||
-p 8000:8000 \
|
||||
-v routstr-data:/app/data \
|
||||
ghcr.io/routstr/proxy:latest
|
||||
# Initial Admin Password
|
||||
ADMIN_PASSWORD=mysecretpassword
|
||||
|
||||
# Node Identity
|
||||
NAME="My AI Node"
|
||||
DESCRIPTION="Fast access to models"
|
||||
NSEC=yournsec
|
||||
|
||||
# Lightning Payouts
|
||||
RECEIVE_LN_ADDRESS=yourname@wallet.com
|
||||
|
||||
```
|
||||
|
||||
## 2. Start the Node
|
||||
|
||||
The recommended way to run Routstr is using Docker Compose, which handles the node, the UI, and optional services like Tor.
|
||||
|
||||
```bash
|
||||
docker compose up -d
|
||||
```
|
||||
|
||||
Verify it's running:
|
||||
@@ -40,14 +56,23 @@ Verify it's running:
|
||||
curl http://localhost:8000/v1/info
|
||||
```
|
||||
|
||||
### Build from Source (Optional)
|
||||
|
||||
If you've cloned the repository and want to build the images yourself:
|
||||
|
||||
```bash
|
||||
docker compose build
|
||||
docker compose up -d
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 2. Configure via Dashboard
|
||||
## 3. Configure via Dashboard
|
||||
|
||||
Open the **Admin Dashboard** at [http://localhost:8000/admin/](http://localhost:8000/admin/).
|
||||
|
||||
!!! note "Default Access"
|
||||
The dashboard has no password by default. Set one immediately in Settings for production use.
|
||||
!!! note "Login"
|
||||
Use the `ADMIN_PASSWORD` you defined in your `.env` file to log in. If you didn't set one, the dashboard will prompt you to set one on first visit.
|
||||
|
||||
### Connect Your AI Providers
|
||||
|
||||
|
||||
45
examples/create_child_keys.py
Normal file
45
examples/create_child_keys.py
Normal file
@@ -0,0 +1,45 @@
|
||||
import json
|
||||
import sys
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
def create_child_keys(base_url: str, api_key: str, count: int = 3) -> list[str]:
|
||||
headers = {"Authorization": f"Bearer {api_key}"}
|
||||
|
||||
print(f"Requesting {count} child keys from {base_url}...")
|
||||
|
||||
child_keys = []
|
||||
|
||||
for i in range(count):
|
||||
try:
|
||||
response = httpx.post(f"{base_url}/v1/balance/child-key", headers=headers)
|
||||
if response.status_code == 200:
|
||||
data = response.json()
|
||||
child_keys.append(data["api_key"])
|
||||
print(
|
||||
f" [{i + 1}] Created: {data['api_key']} (Cost: {data['cost_msats']} msats)"
|
||||
)
|
||||
else:
|
||||
print(f" [{i + 1}] Failed: {response.status_code} - {response.text}")
|
||||
except Exception as e:
|
||||
print(f" [{i + 1}] Error: {str(e)}")
|
||||
|
||||
return child_keys
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
if len(sys.argv) < 2:
|
||||
print("Usage: python create_child_keys.py <api_key_or_cashu_token> [base_url]")
|
||||
sys.exit(1)
|
||||
|
||||
auth_key = sys.argv[1]
|
||||
base_url = sys.argv[2] if len(sys.argv) > 2 else "http://localhost:8000"
|
||||
|
||||
keys = create_child_keys(base_url, auth_key)
|
||||
|
||||
if keys:
|
||||
print("\nSuccessfully created child keys:")
|
||||
print(json.dumps(keys, indent=2))
|
||||
else:
|
||||
print("\nNo child keys were created.")
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
|
||||
32
migrations/versions/02650cd6f028_add_routstr_fees_table.py
Normal file
32
migrations/versions/02650cd6f028_add_routstr_fees_table.py
Normal file
@@ -0,0 +1,32 @@
|
||||
"""add routstr_fees table
|
||||
|
||||
Revision ID: 02650cd6f028
|
||||
Revises: c3d4e5f6a7b8
|
||||
Create Date: 2026-04-24 00:00:00.000000
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "02650cd6f028"
|
||||
down_revision = "c3d4e5f6a7b8"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"routstr_fees",
|
||||
sa.Column("id", sa.Integer(), nullable=False),
|
||||
sa.Column("accumulated_msats", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("total_paid_msats", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("last_paid_at", sa.Integer(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
# Seed with a single row
|
||||
op.execute("INSERT INTO routstr_fees (id, accumulated_msats, total_paid_msats) VALUES (1, 0, 0)")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("routstr_fees")
|
||||
@@ -0,0 +1,37 @@
|
||||
"""add key management and reset fields to api_keys
|
||||
|
||||
Revision ID: 06f81c0fc88d
|
||||
Revises: c2d3e4f5a6b7
|
||||
Create Date: 2026-02-04 22:44:03.311983
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
import sqlmodel
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "06f81c0fc88d"
|
||||
down_revision = "c2d3e4f5a6b7"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column("api_keys", sa.Column("balance_limit", sa.Integer(), nullable=True))
|
||||
op.add_column(
|
||||
"api_keys",
|
||||
sa.Column(
|
||||
"balance_limit_reset", sqlmodel.sql.sqltypes.AutoString(), nullable=True
|
||||
),
|
||||
)
|
||||
op.add_column(
|
||||
"api_keys", sa.Column("balance_limit_reset_date", sa.Integer(), nullable=True)
|
||||
)
|
||||
op.add_column("api_keys", sa.Column("validity_date", sa.Integer(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("api_keys", "validity_date")
|
||||
op.drop_column("api_keys", "balance_limit_reset_date")
|
||||
op.drop_column("api_keys", "balance_limit_reset")
|
||||
op.drop_column("api_keys", "balance_limit")
|
||||
@@ -0,0 +1,32 @@
|
||||
"""add provider_settings to upstream_providers
|
||||
|
||||
Revision ID: 614c0a740e68
|
||||
Revises: 06f81c0fc88d
|
||||
Create Date: 2026-02-13 22:36:53.608737
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "614c0a740e68"
|
||||
down_revision = "06f81c0fc88d"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Check if column exists before adding it
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
columns = [c["name"] for c in inspector.get_columns("upstream_providers")]
|
||||
|
||||
if "provider_settings" not in columns:
|
||||
op.add_column(
|
||||
"upstream_providers",
|
||||
sa.Column("provider_settings", sa.Text(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("upstream_providers", "provider_settings")
|
||||
42
migrations/versions/a776ca70e5fe_add_cashu_refunds_table.py
Normal file
42
migrations/versions/a776ca70e5fe_add_cashu_refunds_table.py
Normal file
@@ -0,0 +1,42 @@
|
||||
"""add cashu_transactions table
|
||||
|
||||
Revision ID: a776ca70e5fe
|
||||
Revises: 614c0a740e68
|
||||
Create Date: 2026-03-11 22:00:01.554762
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
import sqlmodel
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "a776ca70e5fe"
|
||||
down_revision = "614c0a740e68"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"cashu_transactions",
|
||||
sa.Column("id", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||
sa.Column("token", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||
sa.Column("amount", sa.Integer(), nullable=False),
|
||||
sa.Column("unit", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||
sa.Column("mint_url", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
|
||||
sa.Column(
|
||||
"type",
|
||||
sqlmodel.sql.sqltypes.AutoString(),
|
||||
nullable=False,
|
||||
server_default="out",
|
||||
),
|
||||
sa.Column("request_id", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
|
||||
sa.Column("created_at", sa.Integer(), nullable=False),
|
||||
sa.Column("collected", sa.Boolean(), nullable=False),
|
||||
sa.Column("swept", sa.Boolean(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("cashu_transactions")
|
||||
42
migrations/versions/a86e5348850b_.py
Normal file
42
migrations/versions/a86e5348850b_.py
Normal file
@@ -0,0 +1,42 @@
|
||||
"""
|
||||
|
||||
Revision ID: a86e5348850b
|
||||
Revises: b9667ffc5701
|
||||
Create Date: 2026-01-10 18:57:48.475781
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
import sqlmodel
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "a86e5348850b"
|
||||
down_revision = "b9667ffc5701"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Use batch_alter_table for SQLite compatibility
|
||||
with op.batch_alter_table("api_keys", schema=None) as batch_op:
|
||||
batch_op.add_column(
|
||||
sa.Column(
|
||||
"parent_key_hash", sqlmodel.sql.sqltypes.AutoString(), nullable=True
|
||||
)
|
||||
)
|
||||
batch_op.create_index(
|
||||
batch_op.f("ix_api_keys_parent_key_hash"), ["parent_key_hash"], unique=False
|
||||
)
|
||||
batch_op.create_foreign_key(
|
||||
"fk_api_keys_parent_key_hash",
|
||||
"api_keys",
|
||||
["parent_key_hash"],
|
||||
["hashed_key"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
with op.batch_alter_table("api_keys", schema=None) as batch_op:
|
||||
batch_op.drop_constraint("fk_api_keys_parent_key_hash", type_="foreignkey")
|
||||
batch_op.drop_index(batch_op.f("ix_api_keys_parent_key_hash"))
|
||||
batch_op.drop_column("parent_key_hash")
|
||||
@@ -0,0 +1,33 @@
|
||||
"""add forwarded_model_id to models
|
||||
|
||||
Revision ID: b1c2d3e4f5a6
|
||||
Revises: a776ca70e5fe
|
||||
Create Date: 2026-04-05 00:00:00.000000
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
import sqlmodel
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "b1c2d3e4f5a6"
|
||||
down_revision = "a776ca70e5fe"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"models",
|
||||
sa.Column(
|
||||
"forwarded_model_id",
|
||||
sqlmodel.sql.sqltypes.AutoString(),
|
||||
nullable=True,
|
||||
),
|
||||
)
|
||||
# Backfill: set forwarded_model_id = id for all existing rows
|
||||
op.execute("UPDATE models SET forwarded_model_id = id WHERE forwarded_model_id IS NULL")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("models", "forwarded_model_id")
|
||||
@@ -0,0 +1,118 @@
|
||||
"""make upstream provider base_url + api_key unique
|
||||
|
||||
Revision ID: c2d3e4f5a6b7
|
||||
Revises: a86e5348850b
|
||||
Create Date: 2026-01-25 00:00:00.000000
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "c2d3e4f5a6b7"
|
||||
down_revision = "a86e5348850b"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _recreate_table_sqlite(add_base_url_unique: bool) -> None:
|
||||
conn = op.get_bind()
|
||||
existing_tables = {
|
||||
row[0]
|
||||
for row in conn.exec_driver_sql(
|
||||
"SELECT name FROM sqlite_master WHERE type='table'"
|
||||
).fetchall()
|
||||
}
|
||||
if "upstream_providers_old" in existing_tables:
|
||||
if "upstream_providers" in existing_tables:
|
||||
op.drop_table("upstream_providers_old")
|
||||
else:
|
||||
op.execute(
|
||||
"ALTER TABLE upstream_providers_old RENAME TO upstream_providers"
|
||||
)
|
||||
existing_tables.add("upstream_providers")
|
||||
if "upstream_providers" not in existing_tables:
|
||||
return
|
||||
|
||||
constraints = [
|
||||
sa.UniqueConstraint(
|
||||
"base_url",
|
||||
"api_key",
|
||||
name="uq_upstream_providers_base_url_api_key",
|
||||
)
|
||||
]
|
||||
if add_base_url_unique:
|
||||
constraints.append(
|
||||
sa.UniqueConstraint("base_url", name="uq_upstream_providers_base_url")
|
||||
)
|
||||
|
||||
op.execute("ALTER TABLE upstream_providers RENAME TO upstream_providers_old")
|
||||
op.create_table(
|
||||
"upstream_providers",
|
||||
sa.Column(
|
||||
"id", sa.Integer(), primary_key=True, nullable=False, autoincrement=True
|
||||
),
|
||||
sa.Column("provider_type", sa.String(), nullable=False),
|
||||
sa.Column("base_url", sa.String(), nullable=False),
|
||||
sa.Column("api_key", sa.String(), nullable=False),
|
||||
sa.Column("api_version", sa.String(), nullable=True),
|
||||
sa.Column("enabled", sa.Boolean(), nullable=False),
|
||||
sa.Column("provider_fee", sa.Float(), nullable=False, server_default="1.01"),
|
||||
*constraints,
|
||||
)
|
||||
op.execute(
|
||||
"INSERT INTO upstream_providers (id, provider_type, base_url, api_key, api_version, enabled, provider_fee) "
|
||||
"SELECT id, provider_type, base_url, api_key, api_version, enabled, provider_fee "
|
||||
"FROM upstream_providers_old"
|
||||
)
|
||||
op.drop_table("upstream_providers_old")
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
if conn.dialect.name == "sqlite":
|
||||
_recreate_table_sqlite(add_base_url_unique=False)
|
||||
return
|
||||
|
||||
inspector = sa.inspect(conn)
|
||||
for constraint in inspector.get_unique_constraints("upstream_providers"):
|
||||
name = constraint.get("name")
|
||||
if constraint.get("column_names") == ["base_url"] and name:
|
||||
op.drop_constraint(
|
||||
name,
|
||||
"upstream_providers",
|
||||
type_="unique",
|
||||
)
|
||||
index_names = {idx["name"] for idx in inspector.get_indexes("upstream_providers")}
|
||||
if "ix_upstream_providers_base_url" in index_names:
|
||||
op.drop_index("ix_upstream_providers_base_url", table_name="upstream_providers")
|
||||
op.create_unique_constraint(
|
||||
"uq_upstream_providers_base_url_api_key",
|
||||
"upstream_providers",
|
||||
["base_url", "api_key"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
if conn.dialect.name == "sqlite":
|
||||
_recreate_table_sqlite(add_base_url_unique=True)
|
||||
return
|
||||
|
||||
op.drop_constraint(
|
||||
"uq_upstream_providers_base_url_api_key",
|
||||
"upstream_providers",
|
||||
type_="unique",
|
||||
)
|
||||
op.create_unique_constraint(
|
||||
"uq_upstream_providers_base_url",
|
||||
"upstream_providers",
|
||||
["base_url"],
|
||||
)
|
||||
op.create_index(
|
||||
"ix_upstream_providers_base_url",
|
||||
"upstream_providers",
|
||||
["base_url"],
|
||||
unique=True,
|
||||
)
|
||||
@@ -0,0 +1,36 @@
|
||||
"""add source to cashu_transactions
|
||||
|
||||
Revision ID: c3d4e5f6a7b8
|
||||
Revises: b1c2d3e4f5a6
|
||||
Create Date: 2026-04-10 00:00:00.000000
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
import sqlmodel
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "c3d4e5f6a7b8"
|
||||
down_revision = "b1c2d3e4f5a6"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
columns = [col["name"] for col in inspector.get_columns("cashu_transactions")]
|
||||
if "source" not in columns:
|
||||
op.add_column(
|
||||
"cashu_transactions",
|
||||
sa.Column(
|
||||
"source",
|
||||
sqlmodel.sql.sqltypes.AutoString(),
|
||||
nullable=False,
|
||||
server_default="x-cashu",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("cashu_transactions", "source")
|
||||
34
migrations/versions/cli_tokens_table.py
Normal file
34
migrations/versions/cli_tokens_table.py
Normal file
@@ -0,0 +1,34 @@
|
||||
"""add cli_tokens table
|
||||
|
||||
Revision ID: cli_tokens_001
|
||||
Revises: e8f9a0b1c2d3
|
||||
Create Date: 2026-04-25 00:00:00.000000
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "cli_tokens_001"
|
||||
down_revision = "e8f9a0b1c2d3"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"cli_tokens",
|
||||
sa.Column("id", sa.String(), primary_key=True, nullable=False),
|
||||
sa.Column("token", sa.String(), nullable=False, unique=True),
|
||||
sa.Column("name", sa.String(), nullable=False),
|
||||
sa.Column("created_at", sa.Integer(), nullable=False),
|
||||
sa.Column("last_used_at", sa.Integer(), nullable=True),
|
||||
sa.Column("expires_at", sa.Integer(), nullable=True),
|
||||
)
|
||||
op.create_index("ix_cli_tokens_token", "cli_tokens", ["token"], unique=True)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_cli_tokens_token", table_name="cli_tokens")
|
||||
op.drop_table("cli_tokens")
|
||||
@@ -0,0 +1,46 @@
|
||||
"""add api key link to cashu_transactions
|
||||
|
||||
Revision ID: d4e5f6a7b8c9
|
||||
Revises: c3d4e5f6a7b8
|
||||
Create Date: 2026-04-20 00:00:00.000000
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
import sqlmodel
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "d4e5f6a7b8c9"
|
||||
down_revision = "c3d4e5f6a7b8"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
columns = [col["name"] for col in inspector.get_columns("cashu_transactions")]
|
||||
indexes = {index["name"] for index in inspector.get_indexes("cashu_transactions")}
|
||||
|
||||
if "api_key_hashed_key" not in columns:
|
||||
op.add_column(
|
||||
"cashu_transactions",
|
||||
sa.Column(
|
||||
"api_key_hashed_key",
|
||||
sqlmodel.sql.sqltypes.AutoString(),
|
||||
nullable=True,
|
||||
),
|
||||
)
|
||||
|
||||
if "ix_cashu_transactions_api_key_hashed_key" not in indexes:
|
||||
op.create_index(
|
||||
"ix_cashu_transactions_api_key_hashed_key",
|
||||
"cashu_transactions",
|
||||
["api_key_hashed_key"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_cashu_transactions_api_key_hashed_key", table_name="cashu_transactions")
|
||||
op.drop_column("cashu_transactions", "api_key_hashed_key")
|
||||
20
migrations/versions/e8f9a0b1c2d3_merge_heads.py
Normal file
20
migrations/versions/e8f9a0b1c2d3_merge_heads.py
Normal file
@@ -0,0 +1,20 @@
|
||||
"""merge heads: routstr_fees + api_key_to_cashu_transactions
|
||||
|
||||
Revision ID: e8f9a0b1c2d3
|
||||
Revises: 02650cd6f028, d4e5f6a7b8c9
|
||||
Create Date: 2026-04-24 00:00:00.000000
|
||||
"""
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "e8f9a0b1c2d3"
|
||||
down_revision = ("02650cd6f028", "d4e5f6a7b8c9")
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
pass
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
pass
|
||||
355
plans/other-mints-balance-tracking-plan.md
Normal file
355
plans/other-mints-balance-tracking-plan.md
Normal file
@@ -0,0 +1,355 @@
|
||||
# Plan: Track unsupported incoming mints in `other_mints` and include them in balances
|
||||
|
||||
## Goal
|
||||
|
||||
When a Cashu token arrives from a mint that is **not** in `settings.cashu_mints`, we currently create/use a wallet for that mint and swap value into the primary mint. Any change/surplus left behind after `melt()` remains in the foreign `token_wallet`, but that balance is not surfaced by `fetch_all_balances()` because it only iterates over configured mints.
|
||||
|
||||
This plan adds persistent tracking for those foreign mints in a new database table called `other_mints`, and updates balance reporting to include them.
|
||||
|
||||
---
|
||||
|
||||
## Current behavior
|
||||
|
||||
### Incoming unsupported mint flow
|
||||
|
||||
In `routstr/wallet.py`:
|
||||
|
||||
- `recieve_token()` deserializes the token
|
||||
- if `token_obj.mint not in settings.cashu_mints`, it calls `swap_to_primary_mint(token_obj, wallet)`
|
||||
- `swap_to_primary_mint()` calls `token_wallet.melt(...)` to pay the primary mint invoice
|
||||
|
||||
### Important detail: change is retained, not discarded
|
||||
|
||||
The underlying Cashu wallet library keeps any melt change:
|
||||
|
||||
- `Wallet.melt()` constructs blank outputs for change
|
||||
- when the melt succeeds, returned change is reconstructed into proofs
|
||||
- those proofs are appended to `self.proofs` and stored in the wallet DB
|
||||
|
||||
So surplus from unsupported mints is **not discarded**, but it may become invisible operationally.
|
||||
|
||||
### Visibility problem
|
||||
|
||||
`fetch_all_balances()` currently only loops over:
|
||||
|
||||
- `settings.cashu_mints`
|
||||
- units `sat` and `msat`
|
||||
|
||||
This means balances left on unsupported mints are not shown in admin balance reporting.
|
||||
|
||||
---
|
||||
|
||||
## Proposed design
|
||||
|
||||
## 1. Add a new DB table: `other_mints`
|
||||
|
||||
Add a small table in `routstr/core/db.py` to persist unsupported mints we have seen in incoming tokens.
|
||||
|
||||
Suggested schema:
|
||||
|
||||
- `mint_url: str` primary key
|
||||
- `created_at: int`
|
||||
- `last_seen_at: int`
|
||||
|
||||
Minimal model:
|
||||
|
||||
```python
|
||||
class OtherMint(SQLModel, table=True):
|
||||
__tablename__ = "other_mints"
|
||||
|
||||
mint_url: str = Field(primary_key=True)
|
||||
created_at: int = Field(default_factory=lambda: int(time.time()))
|
||||
last_seen_at: int = Field(default_factory=lambda: int(time.time()))
|
||||
```
|
||||
|
||||
Why minimal:
|
||||
|
||||
- the only required function is mint discovery/tracking
|
||||
- unit handling can remain dynamic via existing balance queries over `sat` and `msat`
|
||||
|
||||
---
|
||||
|
||||
## 2. Add DB helpers for `other_mints`
|
||||
|
||||
In `routstr/core/db.py`, add helper functions:
|
||||
|
||||
### `register_other_mint(mint_url: str) -> None`
|
||||
|
||||
Behavior:
|
||||
|
||||
- if the mint is not present, insert it
|
||||
- if it already exists, update `last_seen_at`
|
||||
|
||||
### `list_other_mints(session) -> list[str]`
|
||||
|
||||
Behavior:
|
||||
|
||||
- return all tracked unsupported mint URLs
|
||||
|
||||
Optional later:
|
||||
|
||||
- `delete_other_mint(...)`
|
||||
- admin cleanup helpers
|
||||
|
||||
---
|
||||
|
||||
## 3. Register unsupported mints during token receipt
|
||||
|
||||
Update `recieve_token()` in `routstr/wallet.py`.
|
||||
|
||||
Current logic:
|
||||
|
||||
```python
|
||||
if token_obj.mint not in settings.cashu_mints:
|
||||
return await swap_to_primary_mint(token_obj, wallet)
|
||||
```
|
||||
|
||||
Planned logic:
|
||||
|
||||
```python
|
||||
if token_obj.mint not in settings.cashu_mints:
|
||||
await db.register_other_mint(token_obj.mint)
|
||||
return await swap_to_primary_mint(token_obj, wallet)
|
||||
```
|
||||
|
||||
Why here:
|
||||
|
||||
- this is the earliest reliable point where we know the mint came in via an actual token
|
||||
- this is exactly the path that can leave foreign-mint change behind
|
||||
- it avoids needing to infer unsupported mints later from wallet internals
|
||||
|
||||
---
|
||||
|
||||
## 4. Update `fetch_all_balances()` to include `other_mints`
|
||||
|
||||
Current behavior only includes configured mints.
|
||||
|
||||
Planned behavior:
|
||||
|
||||
- load tracked unsupported mints from DB
|
||||
- combine them with `settings.cashu_mints`
|
||||
- dedupe while preserving order
|
||||
- fetch balances for all tracked mints across requested units
|
||||
|
||||
Conceptual flow:
|
||||
|
||||
```python
|
||||
tracked_mints = dedupe(settings.cashu_mints + other_mints_from_db)
|
||||
```
|
||||
|
||||
Then existing per-mint/per-unit balance logic can remain mostly unchanged.
|
||||
|
||||
This ensures that retained change on unsupported mints becomes visible in admin balance reporting.
|
||||
|
||||
---
|
||||
|
||||
## 5. Add a balance source marker
|
||||
|
||||
Extend `BalanceDetail` in `routstr/wallet.py` to identify whether a balance row comes from a configured mint or an `other_mints` entry.
|
||||
|
||||
Suggested field:
|
||||
|
||||
- `source: str` with values:
|
||||
- `"configured"`
|
||||
- `"other"`
|
||||
|
||||
Updated shape:
|
||||
|
||||
```python
|
||||
class BalanceDetail(TypedDict, total=False):
|
||||
mint_url: str
|
||||
unit: str
|
||||
source: str
|
||||
wallet_balance: int
|
||||
user_balance: int
|
||||
owner_balance: int
|
||||
error: str
|
||||
```
|
||||
|
||||
Why this helps:
|
||||
|
||||
- admin can distinguish normal configured wallet balances from foreign/unsupported balances
|
||||
- avoids confusion if unexpected mint URLs show up in the balances API/UI
|
||||
|
||||
---
|
||||
|
||||
## 6. Admin/API impact
|
||||
|
||||
Backend impact is minimal because `/admin/api/balances` already returns `fetch_all_balances()` output.
|
||||
|
||||
Effects:
|
||||
|
||||
- supported mints continue to show as before
|
||||
- tracked unsupported mints will also appear
|
||||
- UI can optionally display the new `source` field
|
||||
|
||||
No API contract break is expected if the frontend ignores unknown fields.
|
||||
|
||||
---
|
||||
|
||||
## 7. Payout behavior: do not change in phase 1
|
||||
|
||||
`periodic_payout()` currently only iterates over `settings.cashu_mints`.
|
||||
|
||||
Recommendation for this change:
|
||||
|
||||
- **do not** expand `periodic_payout()` to include `other_mints` yet
|
||||
- only improve visibility through balance reporting
|
||||
|
||||
Reason:
|
||||
|
||||
- automatic payout from unsupported/foreign mints may be operationally undesirable
|
||||
- visibility should come first, automation second
|
||||
|
||||
Possible future phase:
|
||||
|
||||
- add optional sweeping/payout support for `other_mints`
|
||||
- or provide an admin-triggered withdrawal/sweep flow
|
||||
|
||||
---
|
||||
|
||||
## 8. Logging improvements (optional)
|
||||
|
||||
Optional follow-up improvement in `swap_to_primary_mint()`:
|
||||
|
||||
- capture the return value from `token_wallet.melt(...)`
|
||||
- if feasible, log any reported change amount
|
||||
- otherwise, rely on wallet balance reporting to surface residual amounts
|
||||
|
||||
This is useful but not required for the first implementation.
|
||||
|
||||
---
|
||||
|
||||
## Files to change
|
||||
|
||||
### `routstr/core/db.py`
|
||||
|
||||
Add:
|
||||
|
||||
- `OtherMint` SQLModel
|
||||
- `register_other_mint()`
|
||||
- `list_other_mints()`
|
||||
|
||||
### `migrations/versions/<new_revision>_add_other_mints_table.py`
|
||||
|
||||
Create migration to add the `other_mints` table.
|
||||
|
||||
### `routstr/wallet.py`
|
||||
|
||||
Update:
|
||||
|
||||
- `recieve_token()` to register unsupported mints
|
||||
- `BalanceDetail` to include `source`
|
||||
- `fetch_all_balances()` to include both configured and tracked unsupported mints
|
||||
|
||||
### `routstr/core/admin.py`
|
||||
|
||||
Likely no backend changes required unless a dedicated `other_mints` API is desired.
|
||||
|
||||
---
|
||||
|
||||
## Behavior rules
|
||||
|
||||
### Register a mint when
|
||||
|
||||
- an incoming token is processed
|
||||
- the token mint is not in `settings.cashu_mints`
|
||||
|
||||
### Do not remove automatically when
|
||||
|
||||
- balance reaches zero
|
||||
|
||||
Reason:
|
||||
|
||||
- historical visibility is useful
|
||||
- avoids flapping entries in the admin balance list
|
||||
- mint may receive additional unsupported tokens later
|
||||
|
||||
Potential future enhancement:
|
||||
|
||||
- admin endpoint to prune zero-balance `other_mints`
|
||||
|
||||
---
|
||||
|
||||
## Edge cases
|
||||
|
||||
### A mint later becomes configured
|
||||
|
||||
If a mint in `other_mints` is later added to `settings.cashu_mints`:
|
||||
|
||||
- deduplication prevents duplicate balance rows
|
||||
- `source` should resolve to `configured`
|
||||
|
||||
### Unsupported mint with zero balance
|
||||
|
||||
A tracked unsupported mint may show zero balances.
|
||||
|
||||
Initial recommendation:
|
||||
|
||||
- allow it to appear
|
||||
- consider later filtering zero-balance `other` rows if the UI becomes noisy
|
||||
|
||||
### Units
|
||||
|
||||
Balance fetching can continue to query both `sat` and `msat` for each tracked mint.
|
||||
|
||||
If a mint has no proofs in one unit, current error/zero handling can continue to apply.
|
||||
|
||||
---
|
||||
|
||||
## Test plan
|
||||
|
||||
### DB tests
|
||||
|
||||
- registering a new unsupported mint inserts a row
|
||||
- registering the same mint again updates `last_seen_at` without duplication
|
||||
- listing other mints returns expected mint URLs
|
||||
|
||||
### Wallet tests
|
||||
|
||||
#### `recieve_token()`
|
||||
|
||||
- when mint is unsupported, `db.register_other_mint()` is called before swap
|
||||
- when mint is configured, `db.register_other_mint()` is not called
|
||||
|
||||
#### `fetch_all_balances()`
|
||||
|
||||
- includes configured mints
|
||||
- includes `other_mints` from DB
|
||||
- dedupes if a mint exists in both configured and other lists
|
||||
- sets `source` correctly
|
||||
|
||||
### Regression tests
|
||||
|
||||
- existing trusted mint balance reporting remains unchanged
|
||||
- `/admin/api/balances` continues to work
|
||||
|
||||
---
|
||||
|
||||
## Recommended implementation order
|
||||
|
||||
1. Add `OtherMint` model to `routstr/core/db.py`
|
||||
2. Add Alembic migration for `other_mints`
|
||||
3. Add `register_other_mint()` and `list_other_mints()` helpers
|
||||
4. Update `recieve_token()` to register unsupported mints
|
||||
5. Update `fetch_all_balances()` to union configured + tracked other mints
|
||||
6. Add `source` to `BalanceDetail`
|
||||
7. Add/adjust tests
|
||||
|
||||
---
|
||||
|
||||
## Summary
|
||||
|
||||
This change solves an operational visibility problem:
|
||||
|
||||
- unsupported incoming mints can leave retained change in foreign wallets
|
||||
- those funds are currently preserved but not surfaced in balance reporting
|
||||
- introducing `other_mints` makes those mints discoverable and auditable
|
||||
- expanding `fetch_all_balances()` ensures their balances are visible in admin tooling
|
||||
|
||||
Recommended scope for the first pass:
|
||||
|
||||
- track unsupported mints in DB
|
||||
- include them in balance reporting
|
||||
- mark them as `source="other"`
|
||||
- do not yet change payout/sweeping behavior
|
||||
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "routstr"
|
||||
version = "0.2.2"
|
||||
version = "0.4.3"
|
||||
description = "Payment proxy for your LLM endpoint using cashu and nostr."
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.11"
|
||||
@@ -13,14 +13,14 @@ dependencies = [
|
||||
"greenlet>=3.2.1",
|
||||
"alembic>=1.13",
|
||||
"python-json-logger>=2.0.0",
|
||||
"cashu",
|
||||
"secp256k1",
|
||||
"cashu>=0.20",
|
||||
"marshmallow>=3.13,<4.0",
|
||||
"websockets>=12.0",
|
||||
"nostr>=0.0.2",
|
||||
"mdurl==0.1.2",
|
||||
"pillow>=10",
|
||||
"openai>=1.98.0",
|
||||
"litellm>=1.55.0",
|
||||
]
|
||||
|
||||
[dependency-groups]
|
||||
@@ -86,4 +86,3 @@ disallow_untyped_decorators = true
|
||||
|
||||
[tool.uv.sources]
|
||||
routstr = { workspace = true }
|
||||
secp256k1 = { git = "https://github.com/saschanaz/secp256k1-py", branch = "upgrade060" }
|
||||
|
||||
@@ -84,93 +84,26 @@ def get_provider_penalty(provider: "BaseUpstreamProvider") -> float:
|
||||
return penalty
|
||||
|
||||
|
||||
def should_prefer_model(
|
||||
candidate_model: "Model",
|
||||
candidate_provider: "BaseUpstreamProvider",
|
||||
current_model: "Model",
|
||||
current_provider: "BaseUpstreamProvider",
|
||||
alias: str,
|
||||
) -> bool:
|
||||
"""Determine if candidate model should replace current model for an alias.
|
||||
|
||||
This is the core decision function for model prioritization. It considers:
|
||||
1. Alias matching quality (exact match vs. canonical slug match)
|
||||
2. Model cost (lower is better)
|
||||
3. Provider penalties (e.g., slight preference against OpenRouter)
|
||||
|
||||
Args:
|
||||
candidate_model: The new model being considered
|
||||
candidate_provider: Provider offering the candidate model
|
||||
current_model: The currently selected model for this alias
|
||||
current_provider: Provider offering the current model
|
||||
alias: The model alias being mapped
|
||||
|
||||
Returns:
|
||||
True if candidate should replace current, False otherwise
|
||||
"""
|
||||
|
||||
def get_base_model_id(model_id: str) -> str:
|
||||
"""Get base model ID by removing provider prefix."""
|
||||
return model_id.split("/", 1)[1] if "/" in model_id else model_id
|
||||
|
||||
def alias_priority(model: "Model") -> int:
|
||||
"""Rank how strong the mapping of alias->model is.
|
||||
|
||||
Highest priority when alias exactly equals the model ID without provider prefix.
|
||||
Next when alias equals canonical slug without prefix. Otherwise lowest.
|
||||
"""
|
||||
model_base = get_base_model_id(model.id)
|
||||
if model_base == alias:
|
||||
return 3
|
||||
if model.canonical_slug:
|
||||
canonical_base = get_base_model_id(model.canonical_slug)
|
||||
if canonical_base == alias:
|
||||
return 2
|
||||
return 1
|
||||
|
||||
candidate_alias_priority = alias_priority(candidate_model)
|
||||
current_alias_priority = alias_priority(current_model)
|
||||
|
||||
# If candidate has better alias match, prefer it regardless of cost
|
||||
if candidate_alias_priority > current_alias_priority:
|
||||
return True
|
||||
|
||||
# If current has better alias match, keep it regardless of cost
|
||||
if current_alias_priority > candidate_alias_priority:
|
||||
return False
|
||||
|
||||
# Same alias priority - compare costs
|
||||
candidate_cost = calculate_model_cost_score(candidate_model)
|
||||
current_cost = calculate_model_cost_score(current_model)
|
||||
|
||||
# Apply provider penalties
|
||||
candidate_adjusted = candidate_cost * get_provider_penalty(candidate_provider)
|
||||
current_adjusted = current_cost * get_provider_penalty(current_provider)
|
||||
|
||||
# Prefer lower adjusted cost
|
||||
should_replace = candidate_adjusted < current_adjusted
|
||||
|
||||
return should_replace
|
||||
|
||||
|
||||
def create_model_mappings(
|
||||
upstreams: list["BaseUpstreamProvider"],
|
||||
overrides_by_id: dict[str, tuple],
|
||||
disabled_model_ids: set[str],
|
||||
) -> tuple[dict[str, "Model"], dict[str, "BaseUpstreamProvider"], dict[str, "Model"]]:
|
||||
) -> tuple[
|
||||
dict[str, "Model"], dict[str, list["BaseUpstreamProvider"]], dict[str, "Model"]
|
||||
]:
|
||||
"""Create optimal model mappings based on cost and provider preferences.
|
||||
|
||||
This is the main entry point for the algorithm. It processes all upstream providers
|
||||
and creates three mappings based on cost optimization:
|
||||
|
||||
1. model_instances: alias -> Model (all model aliases mapped to their Model objects)
|
||||
2. provider_map: alias -> UpstreamProvider (which provider to use for each alias)
|
||||
2. provider_map: alias -> List[UpstreamProvider] (sorted list of providers for each alias)
|
||||
3. unique_models: base_id -> Model (unique models without provider prefixes)
|
||||
|
||||
The algorithm:
|
||||
- Processes non-OpenRouter providers first (they're typically cheaper)
|
||||
- Then processes OpenRouter models (they can still win if cheaper)
|
||||
- For each model alias, uses should_prefer_model() to select the best provider
|
||||
- For each model alias, collects all candidates and sorts them by priority and cost.
|
||||
|
||||
Args:
|
||||
upstreams: List of all upstream provider instances
|
||||
@@ -183,15 +116,34 @@ def create_model_mappings(
|
||||
from .payment.models import _row_to_model
|
||||
from .upstream.helpers import resolve_model_alias
|
||||
|
||||
model_instances: dict[str, "Model"] = {}
|
||||
provider_map: dict[str, "BaseUpstreamProvider"] = {}
|
||||
candidates: dict[str, list[tuple["Model", "BaseUpstreamProvider"]]] = {}
|
||||
unique_models: dict[str, "Model"] = {}
|
||||
seen_model_provider: set[tuple[str, str]] = set()
|
||||
|
||||
providers_by_db_id: dict[int, "BaseUpstreamProvider"] = {}
|
||||
for upstream in upstreams:
|
||||
db_id = getattr(upstream, "db_id", None)
|
||||
if isinstance(db_id, int):
|
||||
providers_by_db_id[db_id] = upstream
|
||||
|
||||
# Group upstreams by URL and keep only the one with the lowest fee for each URL
|
||||
upstreams_by_url: dict[str, list["BaseUpstreamProvider"]] = {}
|
||||
for upstream in upstreams:
|
||||
url = getattr(upstream, "base_url", "")
|
||||
if url not in upstreams_by_url:
|
||||
upstreams_by_url[url] = []
|
||||
upstreams_by_url[url].append(upstream)
|
||||
|
||||
filtered_upstreams: list["BaseUpstreamProvider"] = []
|
||||
for providers in upstreams_by_url.values():
|
||||
best_provider = min(providers, key=lambda p: p.provider_fee)
|
||||
filtered_upstreams.append(best_provider)
|
||||
|
||||
# Separate OpenRouter from other providers
|
||||
openrouter: "BaseUpstreamProvider" | None = None
|
||||
other_upstreams: list["BaseUpstreamProvider"] = []
|
||||
|
||||
for upstream in upstreams:
|
||||
for upstream in filtered_upstreams:
|
||||
base_url = getattr(upstream, "base_url", "")
|
||||
if base_url == "https://openrouter.ai/api/v1":
|
||||
openrouter = upstream
|
||||
@@ -202,30 +154,31 @@ def create_model_mappings(
|
||||
"""Get base model ID by removing provider prefix."""
|
||||
return model_id.split("/", 1)[1] if "/" in model_id else model_id
|
||||
|
||||
def _maybe_set_alias(
|
||||
def get_provider_identity(upstream: "BaseUpstreamProvider") -> str:
|
||||
"""Get a stable provider identity used for deduplication."""
|
||||
db_id = getattr(upstream, "db_id", None)
|
||||
if isinstance(db_id, int):
|
||||
return f"db:{db_id}"
|
||||
|
||||
provider_type = str(getattr(upstream, "provider_type", "") or "").lower()
|
||||
base_url = str(getattr(upstream, "base_url", "") or "").lower()
|
||||
return f"{provider_type}|{base_url}"
|
||||
|
||||
def _add_candidate(
|
||||
alias: str, model: "Model", provider: "BaseUpstreamProvider"
|
||||
) -> None:
|
||||
"""Set alias to model/provider if not set or if new model is preferred."""
|
||||
"""Add candidate model/provider for an alias."""
|
||||
alias_lower = alias.lower()
|
||||
existing_model = model_instances.get(alias_lower)
|
||||
if not existing_model:
|
||||
# No existing mapping, set it
|
||||
model_instances[alias_lower] = model
|
||||
provider_map[alias_lower] = provider
|
||||
else:
|
||||
# Check if candidate should replace existing
|
||||
existing_provider = provider_map[alias_lower]
|
||||
if should_prefer_model(
|
||||
model, provider, existing_model, existing_provider, alias
|
||||
):
|
||||
model_instances[alias_lower] = model
|
||||
provider_map[alias_lower] = provider
|
||||
if alias_lower not in candidates:
|
||||
candidates[alias_lower] = []
|
||||
candidates[alias_lower].append((model, provider))
|
||||
|
||||
def process_provider_models(
|
||||
upstream: "BaseUpstreamProvider", is_openrouter: bool = False
|
||||
) -> None:
|
||||
"""Process all models from a given provider."""
|
||||
upstream_prefix = getattr(upstream, "upstream_name", None)
|
||||
provider_key = get_provider_identity(upstream)
|
||||
|
||||
for model in upstream.get_cached_models():
|
||||
if not model.enabled or model.id in disabled_model_ids:
|
||||
@@ -242,14 +195,15 @@ def create_model_mappings(
|
||||
|
||||
# Add to unique models
|
||||
base_id = get_base_model_id(model_to_use.id)
|
||||
if not is_openrouter or base_id not in unique_models:
|
||||
unique_key = model_to_use.forwarded_model_id or base_id
|
||||
if not is_openrouter or unique_key not in unique_models:
|
||||
unique_model = model_to_use.copy(
|
||||
update={
|
||||
"id": base_id,
|
||||
"upstream_provider_id": upstream.provider_type,
|
||||
}
|
||||
)
|
||||
unique_models[base_id] = unique_model
|
||||
unique_models[unique_key] = unique_model
|
||||
|
||||
# Get all aliases for this model
|
||||
aliases = resolve_model_alias(
|
||||
@@ -264,23 +218,165 @@ def create_model_mappings(
|
||||
if prefixed_id not in aliases:
|
||||
aliases.append(prefixed_id)
|
||||
|
||||
# Register forwarded_model_id as a routable alias
|
||||
if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases:
|
||||
aliases.append(model_to_use.forwarded_model_id)
|
||||
|
||||
# Try to set each alias
|
||||
for alias in aliases:
|
||||
_maybe_set_alias(alias, model_to_use, upstream)
|
||||
_add_candidate(alias, model_to_use, upstream)
|
||||
seen_model_provider.add((model_to_use.id.lower(), provider_key))
|
||||
|
||||
# Process non-OpenRouter providers first (they're typically cheaper)
|
||||
# Process non-OpenRouter providers first
|
||||
for upstream in other_upstreams:
|
||||
process_provider_models(upstream, is_openrouter=False)
|
||||
|
||||
# Process OpenRouter last - models only win if they're cheaper or better matched
|
||||
# Process OpenRouter last
|
||||
if openrouter:
|
||||
process_provider_models(openrouter, is_openrouter=True)
|
||||
|
||||
# Log provider distribution
|
||||
# Include enabled DB overrides even when provider discovery misses models.
|
||||
# This is important for deployment-based providers like Azure.
|
||||
for model_id, override_data in overrides_by_id.items():
|
||||
if model_id in disabled_model_ids:
|
||||
continue
|
||||
override_row, provider_fee = override_data
|
||||
upstream_provider_id = getattr(override_row, "upstream_provider_id", None)
|
||||
if not isinstance(upstream_provider_id, int):
|
||||
continue
|
||||
|
||||
upstream_for_override = providers_by_db_id.get(upstream_provider_id)
|
||||
if upstream_for_override is None:
|
||||
continue
|
||||
|
||||
provider_key = get_provider_identity(upstream_for_override)
|
||||
dedupe_key = (model_id.lower(), provider_key)
|
||||
if dedupe_key in seen_model_provider:
|
||||
continue
|
||||
|
||||
try:
|
||||
model_to_use = _row_to_model(
|
||||
override_row, apply_provider_fee=True, provider_fee=provider_fee
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Skipping invalid model override while building model mappings",
|
||||
extra={
|
||||
"model_id": model_id,
|
||||
"upstream_provider_id": upstream_provider_id,
|
||||
"error": str(exc),
|
||||
"error_type": type(exc).__name__,
|
||||
},
|
||||
)
|
||||
continue
|
||||
if not model_to_use.enabled:
|
||||
continue
|
||||
|
||||
base_id = get_base_model_id(model_to_use.id)
|
||||
unique_key = model_to_use.forwarded_model_id or base_id
|
||||
is_openrouter = (
|
||||
getattr(upstream_for_override, "base_url", "")
|
||||
== "https://openrouter.ai/api/v1"
|
||||
)
|
||||
if not is_openrouter or unique_key not in unique_models:
|
||||
unique_model = model_to_use.copy(
|
||||
update={
|
||||
"id": base_id,
|
||||
"upstream_provider_id": upstream_for_override.provider_type,
|
||||
}
|
||||
)
|
||||
unique_models[unique_key] = unique_model
|
||||
|
||||
try:
|
||||
aliases = resolve_model_alias(
|
||||
model_to_use.id,
|
||||
model_to_use.canonical_slug,
|
||||
alias_ids=model_to_use.alias_ids,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Skipping model aliases for invalid override model",
|
||||
extra={
|
||||
"model_id": model_id,
|
||||
"upstream_provider_id": upstream_provider_id,
|
||||
"error": str(exc),
|
||||
"error_type": type(exc).__name__,
|
||||
},
|
||||
)
|
||||
continue
|
||||
|
||||
upstream_prefix = getattr(upstream_for_override, "upstream_name", None)
|
||||
if upstream_prefix and "/" not in model_to_use.id:
|
||||
prefixed_id = f"{upstream_prefix}/{model_to_use.id}"
|
||||
if prefixed_id not in aliases:
|
||||
aliases.append(prefixed_id)
|
||||
|
||||
# Register forwarded_model_id as a routable alias
|
||||
if model_to_use.forwarded_model_id and model_to_use.forwarded_model_id not in aliases:
|
||||
aliases.append(model_to_use.forwarded_model_id)
|
||||
|
||||
for alias in aliases:
|
||||
_add_candidate(alias, model_to_use, upstream_for_override)
|
||||
seen_model_provider.add(dedupe_key)
|
||||
|
||||
# Sort candidates and build final maps
|
||||
model_instances: dict[str, "Model"] = {}
|
||||
provider_map: dict[str, list["BaseUpstreamProvider"]] = {}
|
||||
|
||||
def alias_priority(model: "Model", alias: str) -> int:
|
||||
"""Rank how strong the mapping of alias->model is.
|
||||
|
||||
forwarded_model_id is the most specific identifier (set per-provider
|
||||
instance), so a match there should beat a model_id match. This way,
|
||||
when multiple providers have the same model_id but different
|
||||
forwarded_model_ids, the one whose forwarded_model_id equals the
|
||||
requested alias wins.
|
||||
"""
|
||||
if (
|
||||
model.forwarded_model_id
|
||||
and model.forwarded_model_id.lower() == alias
|
||||
):
|
||||
return 5
|
||||
|
||||
if (
|
||||
model.id
|
||||
and model.id.lower() == alias
|
||||
):
|
||||
return 4
|
||||
|
||||
model_base = get_base_model_id(model.id)
|
||||
if model_base == alias:
|
||||
return 3
|
||||
if model.canonical_slug:
|
||||
canonical_base = get_base_model_id(model.canonical_slug)
|
||||
if canonical_base == alias:
|
||||
return 2
|
||||
return 1
|
||||
|
||||
for alias, items in candidates.items():
|
||||
# Sort key: (priority DESC, cost ASC)
|
||||
# Using negative cost for DESC sort overall to keep high priority first
|
||||
def sort_key(item: tuple["Model", "BaseUpstreamProvider"]) -> tuple[int, float]:
|
||||
model, provider = item
|
||||
priority = alias_priority(model, alias)
|
||||
cost = calculate_model_cost_score(model)
|
||||
penalty = get_provider_penalty(provider)
|
||||
adjusted_cost = cost * penalty
|
||||
return (priority, -adjusted_cost)
|
||||
|
||||
items.sort(key=sort_key, reverse=True)
|
||||
|
||||
best_model, best_provider = items[0]
|
||||
model_instances[alias] = best_model
|
||||
provider_map[alias] = [p for _, p in items]
|
||||
|
||||
# Log provider distribution (using top provider for stats)
|
||||
provider_counts: dict[str, int] = {}
|
||||
for provider in provider_map.values():
|
||||
provider_name = getattr(provider, "upstream_name", "unknown")
|
||||
provider_counts[provider_name] = provider_counts.get(provider_name, 0) + 1
|
||||
for providers in provider_map.values():
|
||||
if providers:
|
||||
provider = providers[0]
|
||||
provider_name = getattr(provider, "upstream_name", "unknown")
|
||||
provider_counts[provider_name] = provider_counts.get(provider_name, 0) + 1
|
||||
|
||||
logger.debug(
|
||||
f"Updated model mappings with ({len(unique_models)} unique models and {len(model_instances)} aliases)",
|
||||
|
||||
773
routstr/auth.py
773
routstr/auth.py
File diff suppressed because it is too large
Load Diff
@@ -1,13 +1,22 @@
|
||||
import asyncio
|
||||
import hashlib
|
||||
import time
|
||||
from time import monotonic
|
||||
from typing import Annotated, NoReturn
|
||||
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException
|
||||
from fastapi.responses import JSONResponse
|
||||
from pydantic import BaseModel
|
||||
from sqlmodel import col, select, update
|
||||
|
||||
from .auth import validate_bearer_key
|
||||
from .core.db import ApiKey, AsyncSession, get_session
|
||||
from .auth import get_billing_key, validate_bearer_key
|
||||
from .core.db import (
|
||||
ApiKey,
|
||||
AsyncSession,
|
||||
CashuTransaction,
|
||||
get_session,
|
||||
store_cashu_transaction,
|
||||
)
|
||||
from .core.logging import get_logger
|
||||
from .core.settings import settings
|
||||
from .lightning import lightning_router
|
||||
@@ -32,14 +41,49 @@ async def get_key_from_header(
|
||||
)
|
||||
|
||||
|
||||
async def get_balance_info(key: ApiKey, session: AsyncSession) -> dict:
|
||||
billing_key = await get_billing_key(key, session)
|
||||
info = {
|
||||
"api_key": "sk-" + key.hashed_key,
|
||||
"balance": billing_key.total_balance,
|
||||
"reserved": billing_key.reserved_balance,
|
||||
"is_child": key.parent_key_hash is not None,
|
||||
"parent_key": "sk-" + key.parent_key_hash if key.parent_key_hash else None,
|
||||
"total_requests": key.total_requests,
|
||||
"total_spent": key.total_spent,
|
||||
"balance_limit": key.balance_limit,
|
||||
"balance_limit_reset": key.balance_limit_reset,
|
||||
"validity_date": key.validity_date,
|
||||
}
|
||||
|
||||
if not key.parent_key_hash:
|
||||
# Fetch child keys if this is a parent key
|
||||
statement = select(ApiKey).where(ApiKey.parent_key_hash == key.hashed_key)
|
||||
results = await session.exec(statement)
|
||||
child_keys = results.all()
|
||||
if child_keys:
|
||||
info["child_keys"] = [
|
||||
{
|
||||
"api_key": "sk-" + ck.hashed_key,
|
||||
"total_requests": ck.total_requests,
|
||||
"total_spent": ck.total_spent,
|
||||
"balance_limit": ck.balance_limit,
|
||||
"balance_limit_reset": ck.balance_limit_reset,
|
||||
"validity_date": ck.validity_date,
|
||||
}
|
||||
for ck in child_keys
|
||||
]
|
||||
|
||||
return info
|
||||
|
||||
|
||||
# TODO: remove this endpoint when frontend is updated
|
||||
@router.get("/", include_in_schema=False)
|
||||
async def account_info(key: ApiKey = Depends(get_key_from_header)) -> dict:
|
||||
return {
|
||||
"api_key": "sk-" + key.hashed_key,
|
||||
"balance": key.balance,
|
||||
"reserved": key.reserved_balance,
|
||||
}
|
||||
async def account_info(
|
||||
key: ApiKey = Depends(get_key_from_header),
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> dict:
|
||||
return await get_balance_info(key, session)
|
||||
|
||||
|
||||
# TODO: Implement POST /v1/wallet/create endpoint
|
||||
@@ -56,9 +100,24 @@ async def account_info(key: ApiKey = Depends(get_key_from_header)) -> dict:
|
||||
|
||||
@router.get("/create")
|
||||
async def create_balance(
|
||||
initial_balance_token: str, session: AsyncSession = Depends(get_session)
|
||||
initial_balance_token: str,
|
||||
balance_limit: int | None = None,
|
||||
balance_limit_reset: str | None = None,
|
||||
validity_date: int | None = None,
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> dict:
|
||||
key = await validate_bearer_key(initial_balance_token, session)
|
||||
|
||||
if balance_limit is not None or balance_limit_reset or validity_date:
|
||||
key.balance_limit = balance_limit
|
||||
key.balance_limit_reset = balance_limit_reset
|
||||
key.validity_date = validity_date
|
||||
if balance_limit_reset:
|
||||
key.balance_limit_reset_date = int(time.time())
|
||||
session.add(key)
|
||||
await session.commit()
|
||||
await session.refresh(key)
|
||||
|
||||
return {
|
||||
"api_key": "sk-" + key.hashed_key,
|
||||
"balance": key.balance,
|
||||
@@ -66,12 +125,11 @@ async def create_balance(
|
||||
|
||||
|
||||
@router.get("/info")
|
||||
async def wallet_info(key: ApiKey = Depends(get_key_from_header)) -> dict:
|
||||
return {
|
||||
"api_key": "sk-" + key.hashed_key,
|
||||
"balance": key.balance,
|
||||
"reserved": key.reserved_balance,
|
||||
}
|
||||
async def wallet_info(
|
||||
key: ApiKey = Depends(get_key_from_header),
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> dict:
|
||||
return await get_balance_info(key, session)
|
||||
|
||||
|
||||
class TopupRequest(BaseModel):
|
||||
@@ -85,6 +143,8 @@ async def topup_wallet_endpoint(
|
||||
key: ApiKey = Depends(get_key_from_header),
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> dict[str, int]:
|
||||
billing_key = await get_billing_key(key, session)
|
||||
|
||||
if topup_request is not None:
|
||||
cashu_token = topup_request.cashu_token
|
||||
if cashu_token is None:
|
||||
@@ -94,16 +154,30 @@ async def topup_wallet_endpoint(
|
||||
if len(cashu_token) < 10 or "cashu" not in cashu_token:
|
||||
raise HTTPException(status_code=400, detail="Invalid token format")
|
||||
try:
|
||||
amount_msats = await credit_balance(cashu_token, key, session)
|
||||
amount_msats = await credit_balance(cashu_token, billing_key, session)
|
||||
except ValueError as e:
|
||||
error_msg = str(e)
|
||||
if "already spent" in error_msg.lower():
|
||||
raise HTTPException(status_code=400, detail="Token already spent")
|
||||
elif "invalid" in error_msg.lower() or "decode" in error_msg.lower():
|
||||
raise HTTPException(status_code=400, detail="Invalid token format")
|
||||
elif "insufficient" in error_msg.lower() or "melt fee" in error_msg.lower():
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Token value is too small to cover swap fees. {error_msg}",
|
||||
)
|
||||
elif "failed to melt" in error_msg.lower():
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Failed to swap foreign mint token. {error_msg}",
|
||||
)
|
||||
else:
|
||||
raise HTTPException(status_code=400, detail="Failed to redeem token")
|
||||
except Exception:
|
||||
raise HTTPException(status_code=400, detail=f"Failed to redeem token: {error_msg}")
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"topup_wallet_endpoint: unhandled error",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
)
|
||||
raise HTTPException(status_code=500, detail="Internal server error")
|
||||
return {"msats": amount_msats}
|
||||
|
||||
@@ -137,23 +211,111 @@ async def _refund_cache_set(authorization: str, value: dict[str, str]) -> None:
|
||||
_refund_cache[key] = (expiry, value)
|
||||
|
||||
|
||||
@router.post("/refund")
|
||||
async def _lookup_key_no_create(
|
||||
bearer_value: str, session: AsyncSession
|
||||
) -> ApiKey | None:
|
||||
"""Look up an existing API key without creating one Used by the refund endpoint"""
|
||||
if bearer_value.startswith("sk-"):
|
||||
return await session.get(ApiKey, bearer_value[3:])
|
||||
if bearer_value.startswith("cashu"):
|
||||
hashed = hashlib.sha256(bearer_value.encode()).hexdigest()
|
||||
return await session.get(ApiKey, hashed)
|
||||
return None
|
||||
|
||||
|
||||
async def _restore_balance(
|
||||
session: AsyncSession, hashed_key: str, balance: int, reserved_balance: int, mint_url: str
|
||||
) -> None:
|
||||
"""Restore balance after a failed refund mint attempt."""
|
||||
restore_stmt = (
|
||||
update(ApiKey)
|
||||
.where(col(ApiKey.hashed_key) == hashed_key)
|
||||
.values(
|
||||
balance=col(ApiKey.balance) + balance,
|
||||
reserved_balance=col(ApiKey.reserved_balance) + reserved_balance,
|
||||
)
|
||||
)
|
||||
await session.exec(restore_stmt) # type: ignore[call-overload]
|
||||
await session.commit()
|
||||
logger.info(
|
||||
"refund_wallet_endpoint: balance restored after mint failure",
|
||||
extra={"hashed_key": hashed_key, "restored_balance": balance, "mint_url": mint_url},
|
||||
)
|
||||
|
||||
|
||||
@router.post("/refund", response_model=None)
|
||||
async def refund_wallet_endpoint(
|
||||
authorization: Annotated[str, Header(...)],
|
||||
authorization: Annotated[str | None, Header()] = None,
|
||||
x_cashu: Annotated[str | None, Header()] = None,
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> dict[str, str]:
|
||||
if not authorization.startswith("Bearer "):
|
||||
) -> JSONResponse | dict[str, str]:
|
||||
if x_cashu:
|
||||
# Find the "in" transaction by the original payment token
|
||||
in_tx_result = await session.exec(
|
||||
select(CashuTransaction).where(
|
||||
CashuTransaction.token == x_cashu,
|
||||
CashuTransaction.type == "in",
|
||||
)
|
||||
)
|
||||
in_tx = in_tx_result.first()
|
||||
if in_tx is None:
|
||||
raise HTTPException(status_code=404, detail="Refund not found")
|
||||
|
||||
# Use the request_id to find the associated "out" (refund) transaction
|
||||
if in_tx.request_id is None:
|
||||
raise HTTPException(status_code=404, detail="Refund not found")
|
||||
|
||||
out_tx_result = await session.exec(
|
||||
select(CashuTransaction).where(
|
||||
CashuTransaction.request_id == in_tx.request_id,
|
||||
CashuTransaction.type == "out",
|
||||
)
|
||||
)
|
||||
out_tx = out_tx_result.first()
|
||||
if out_tx is None:
|
||||
raise HTTPException(status_code=404, detail="Refund not found")
|
||||
if out_tx.swept:
|
||||
raise HTTPException(status_code=410, detail="Refund has been swept")
|
||||
|
||||
out_tx.collected = True
|
||||
session.add(out_tx)
|
||||
await session.commit()
|
||||
body: dict[str, str] = {"token": out_tx.token}
|
||||
if out_tx.unit == "sat":
|
||||
body["sats"] = str(out_tx.amount)
|
||||
else:
|
||||
body["msats"] = str(out_tx.amount)
|
||||
return JSONResponse(content=body, headers={"X-Cashu": out_tx.token})
|
||||
|
||||
if authorization is None or not authorization.startswith("Bearer "):
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Invalid authorization. Use 'Bearer <cashu-token>' or 'Bearer <api-key>'",
|
||||
)
|
||||
|
||||
bearer_value: str = authorization[7:]
|
||||
key: ApiKey | None = await _lookup_key_no_create(bearer_value, session)
|
||||
if key is None:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Key not found. Deposit first via /v1/wallet/create before requesting a refund.",
|
||||
)
|
||||
|
||||
if cached := await _refund_cache_get(bearer_value):
|
||||
return cached
|
||||
if key.total_balance <= 0:
|
||||
if cached := await _refund_cache_get(bearer_value):
|
||||
return cached
|
||||
|
||||
key: ApiKey = await validate_bearer_key(bearer_value, session)
|
||||
if key.parent_key_hash:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot refund child key. Please refund the parent key instead.",
|
||||
)
|
||||
|
||||
if key.reserved_balance > 0:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot refund key. There are ongoing requests for this api key.",
|
||||
)
|
||||
|
||||
remaining_balance_msats: int = key.total_balance
|
||||
|
||||
@@ -167,22 +329,51 @@ async def refund_wallet_endpoint(
|
||||
elif remaining_balance <= 0:
|
||||
raise HTTPException(status_code=400, detail="No balance to refund")
|
||||
|
||||
# Perform refund operation first, before modifying balance
|
||||
# Capture values before debit — the session may refresh key after commit
|
||||
pre_debit_balance = key.balance
|
||||
pre_debit_reserved = key.reserved_balance
|
||||
|
||||
# --- DEBIT FIRST: atomically zero the balance before minting tokens ---
|
||||
# This prevents the race where a concurrent topup/spend happens between
|
||||
# reading the balance and minting the refund token (double-spend).
|
||||
debit_stmt = (
|
||||
update(ApiKey)
|
||||
.where(col(ApiKey.hashed_key) == key.hashed_key)
|
||||
.where(col(ApiKey.balance) == pre_debit_balance)
|
||||
.where(col(ApiKey.reserved_balance) == pre_debit_reserved)
|
||||
.values(balance=0, reserved_balance=0)
|
||||
)
|
||||
debit_result = await session.exec(debit_stmt) # type: ignore[call-overload]
|
||||
await session.commit()
|
||||
|
||||
if debit_result.rowcount == 0:
|
||||
# Balance changed between read and debit — another request is active
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail="Balance changed concurrently. Please retry the refund.",
|
||||
)
|
||||
|
||||
# --- MINT: balance is locked at zero, safe to create the refund token ---
|
||||
# Proofs from untrusted mints are swapped to primary_mint on receive.
|
||||
# Use primary_mint unless key.refund_mint_url is an explicitly trusted mint.
|
||||
effective_refund_mint = (
|
||||
key.refund_mint_url
|
||||
if key.refund_mint_url and key.refund_mint_url in settings.cashu_mints
|
||||
else settings.primary_mint
|
||||
)
|
||||
try:
|
||||
if key.refund_address:
|
||||
from .core.settings import settings as global_settings
|
||||
|
||||
await send_to_lnurl(
|
||||
remaining_balance,
|
||||
key.refund_currency or "sat",
|
||||
key.refund_mint_url or global_settings.primary_mint,
|
||||
effective_refund_mint,
|
||||
key.refund_address,
|
||||
)
|
||||
result = {"recipient": key.refund_address}
|
||||
else:
|
||||
refund_currency = key.refund_currency or "sat"
|
||||
token = await send_token(
|
||||
remaining_balance, refund_currency, key.refund_mint_url
|
||||
remaining_balance, refund_currency, effective_refund_mint
|
||||
)
|
||||
result = {"token": token}
|
||||
|
||||
@@ -191,30 +382,109 @@ async def refund_wallet_endpoint(
|
||||
else:
|
||||
result["msats"] = str(remaining_balance_msats)
|
||||
|
||||
if "token" in result:
|
||||
logger.info(
|
||||
"refund_wallet_endpoint: cashu token issued",
|
||||
extra={
|
||||
"path": "/v1/wallet/refund",
|
||||
"token": result["token"],
|
||||
"amount": remaining_balance,
|
||||
"currency": key.refund_currency or "sat",
|
||||
},
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
# Re-raise HTTP exceptions (like 400 for balance too small)
|
||||
# Minting failed — restore the debited balance
|
||||
await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "")
|
||||
raise
|
||||
except Exception as e:
|
||||
# If refund fails, don't modify the database
|
||||
# Minting failed — restore the debited balance
|
||||
await _restore_balance(session, key.hashed_key, pre_debit_balance, pre_debit_reserved, key.refund_mint_url or "")
|
||||
error_msg = str(e)
|
||||
logger.error(
|
||||
"refund_wallet_endpoint: mint/send failed",
|
||||
extra={
|
||||
"error": error_msg,
|
||||
"error_type": type(e).__name__,
|
||||
"hashed_key": key.hashed_key,
|
||||
"remaining_balance": remaining_balance,
|
||||
"refund_currency": key.refund_currency,
|
||||
"refund_mint_url": key.refund_mint_url,
|
||||
"has_refund_address": bool(key.refund_address),
|
||||
},
|
||||
)
|
||||
if (
|
||||
"mint" in error_msg.lower()
|
||||
or "connection" in error_msg.lower()
|
||||
or isinstance(e, Exception)
|
||||
and "ConnectError" in str(type(e))
|
||||
or "ConnectError" in str(type(e))
|
||||
):
|
||||
raise HTTPException(status_code=503, detail="Mint service unavailable")
|
||||
raise HTTPException(status_code=503, detail=f"Mint service unavailable: {error_msg}")
|
||||
else:
|
||||
raise HTTPException(status_code=500, detail="Refund failed")
|
||||
raise HTTPException(status_code=500, detail=f"Refund failed: {error_msg}")
|
||||
|
||||
await _refund_cache_set(bearer_value, result)
|
||||
|
||||
await session.delete(key)
|
||||
await session.commit()
|
||||
if "token" in result:
|
||||
try:
|
||||
await store_cashu_transaction(
|
||||
token=result["token"],
|
||||
amount=remaining_balance,
|
||||
unit=key.refund_currency or "sat",
|
||||
mint_url=key.refund_mint_url,
|
||||
typ="out",
|
||||
collected=False,
|
||||
source="apikey",
|
||||
api_key_hashed_key=key.hashed_key,
|
||||
)
|
||||
except Exception:
|
||||
pass # store_cashu_transaction already logs
|
||||
|
||||
logger.info(
|
||||
"refund_wallet_endpoint: refund successful",
|
||||
extra={
|
||||
"refunded_msats": remaining_balance_msats,
|
||||
"previous_reserved_balance": key.reserved_balance,
|
||||
},
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
@router.get("/history")
|
||||
async def wallet_history(
|
||||
key: ApiKey = Depends(get_key_from_header),
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> dict[str, list[dict[str, str | int | bool | None]]]:
|
||||
if key.parent_key_hash:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot view child key history. Please use the parent key instead.",
|
||||
)
|
||||
|
||||
result = await session.exec(
|
||||
select(CashuTransaction)
|
||||
.where(CashuTransaction.api_key_hashed_key == key.hashed_key)
|
||||
.order_by(col(CashuTransaction.created_at).desc())
|
||||
)
|
||||
transactions = result.all()
|
||||
return {
|
||||
"transactions": [
|
||||
{
|
||||
"id": tx.id,
|
||||
"type": tx.type,
|
||||
"source": tx.source,
|
||||
"amount": tx.amount,
|
||||
"unit": tx.unit,
|
||||
"mint_url": tx.mint_url,
|
||||
"created_at": tx.created_at,
|
||||
"collected": tx.collected,
|
||||
"swept": tx.swept,
|
||||
}
|
||||
for tx in transactions
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
@router.post("/donate")
|
||||
async def donate(token: str, ref: str | None = None) -> str:
|
||||
try:
|
||||
@@ -228,6 +498,138 @@ async def donate(token: str, ref: str | None = None) -> str:
|
||||
return "Invalid token."
|
||||
|
||||
|
||||
class ChildKeyRequest(BaseModel):
|
||||
count: int
|
||||
balance_limit: int | None = None
|
||||
balance_limit_reset: str | None = None
|
||||
validity_date: int | None = None
|
||||
|
||||
|
||||
@router.post("/child-key")
|
||||
async def create_child_key(
|
||||
payload: ChildKeyRequest,
|
||||
key: ApiKey = Depends(get_key_from_header),
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> dict:
|
||||
"""Creates one or more child API keys that use the parent's balance."""
|
||||
# Log incoming request for debugging
|
||||
logger.debug(f"Child key creation request: count={payload.count}")
|
||||
|
||||
count = payload.count
|
||||
if count < 1 or count > 50:
|
||||
raise HTTPException(status_code=400, detail="Count must be between 1 and 50.")
|
||||
|
||||
# Check if this is already a child key
|
||||
if key.parent_key_hash:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Cannot create a child key for another child key.",
|
||||
)
|
||||
|
||||
cost_per_key = settings.child_key_cost
|
||||
total_cost = cost_per_key * count
|
||||
|
||||
if key.total_balance < total_cost:
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail=f"Insufficient balance to create {count} child keys. {total_cost} mSats required.",
|
||||
)
|
||||
|
||||
# Deduct cost from parent
|
||||
key.balance -= total_cost
|
||||
key.total_spent += total_cost
|
||||
session.add(key)
|
||||
|
||||
# Generate new keys
|
||||
import secrets
|
||||
|
||||
new_keys = []
|
||||
for _ in range(count):
|
||||
new_key_raw = secrets.token_hex(32)
|
||||
new_key_hash = new_key_raw # We use the raw key as the hash for sk- keys
|
||||
|
||||
child_key = ApiKey(
|
||||
hashed_key=new_key_hash,
|
||||
balance=0,
|
||||
parent_key_hash=key.hashed_key,
|
||||
balance_limit=payload.balance_limit,
|
||||
balance_limit_reset=payload.balance_limit_reset,
|
||||
balance_limit_reset_date=int(time.time())
|
||||
if payload.balance_limit_reset
|
||||
else None,
|
||||
validity_date=payload.validity_date,
|
||||
)
|
||||
session.add(child_key)
|
||||
new_keys.append("sk-" + new_key_hash)
|
||||
|
||||
await session.commit()
|
||||
|
||||
response_data = {
|
||||
"api_keys": new_keys,
|
||||
"count": count,
|
||||
"cost_msats": total_cost,
|
||||
"cost_sats": total_cost // 1000,
|
||||
"parent_balance": key.balance,
|
||||
"parent_balance_sats": key.balance // 1000,
|
||||
}
|
||||
logger.debug(f"Child key creation response: {response_data}")
|
||||
return response_data
|
||||
|
||||
|
||||
class ChildKeyResetRequest(BaseModel):
|
||||
child_key: str
|
||||
|
||||
|
||||
@router.post("/child-key/reset")
|
||||
async def reset_child_key_spent(
|
||||
payload: ChildKeyResetRequest,
|
||||
key: ApiKey = Depends(get_key_from_header),
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> dict:
|
||||
"""Resets the total_spent of a child key. Must be called by the parent."""
|
||||
child_key_raw = payload.child_key
|
||||
if child_key_raw.startswith("sk-"):
|
||||
child_key_raw = child_key_raw[3:]
|
||||
|
||||
child_key = await session.get(ApiKey, child_key_raw)
|
||||
if not child_key:
|
||||
raise HTTPException(status_code=404, detail="Child key not found.")
|
||||
|
||||
if child_key.parent_key_hash != key.hashed_key:
|
||||
raise HTTPException(
|
||||
status_code=403, detail="Unauthorized. You are not the parent of this key."
|
||||
)
|
||||
|
||||
child_key.total_spent = 0
|
||||
if child_key.balance_limit_reset:
|
||||
child_key.balance_limit_reset_date = int(time.time())
|
||||
session.add(child_key)
|
||||
await session.commit()
|
||||
|
||||
return {"success": True, "message": "Child key balance reset successfully."}
|
||||
|
||||
|
||||
@router.get("/cashu-refund/{payment_token_hash}")
|
||||
async def get_cashu_refund(
|
||||
payment_token_hash: str,
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> dict:
|
||||
"""Retrieve a stored Cashu refund token by the hash of the original payment token."""
|
||||
result = await session.get(CashuTransaction, payment_token_hash)
|
||||
if result is None:
|
||||
raise HTTPException(status_code=404, detail="Refund not found")
|
||||
if result.swept:
|
||||
raise HTTPException(status_code=410, detail="Refund has been swept")
|
||||
result.collected = True
|
||||
session.add(result)
|
||||
await session.commit()
|
||||
return {
|
||||
"refund_token": result.token,
|
||||
"amount": result.amount,
|
||||
"unit": result.unit,
|
||||
}
|
||||
|
||||
|
||||
@router.api_route(
|
||||
"/{path:path}",
|
||||
methods=["GET", "POST", "PUT", "DELETE"],
|
||||
@@ -240,7 +642,7 @@ async def wallet_catch_all(path: str) -> NoReturn:
|
||||
)
|
||||
|
||||
|
||||
balance_router.include_router(lightning_router)
|
||||
balance_router.include_router(lightning_router, include_in_schema=False)
|
||||
balance_router.include_router(router)
|
||||
|
||||
deprecated_wallet_router = APIRouter(prefix="/v1/wallet", include_in_schema=False)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,12 +1,18 @@
|
||||
import os
|
||||
import pathlib
|
||||
import sqlite3
|
||||
import time
|
||||
import uuid
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import AsyncGenerator
|
||||
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
from alembic.util.exc import CommandError
|
||||
from sqlalchemy import UniqueConstraint
|
||||
from sqlalchemy.exc import OperationalError
|
||||
from sqlalchemy.ext.asyncio.engine import create_async_engine
|
||||
from sqlmodel import Field, Relationship, SQLModel, func, select, update
|
||||
from sqlmodel import Field, Relationship, SQLModel, col, func, select, update
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from .logging import get_logger
|
||||
@@ -47,6 +53,25 @@ class ApiKey(SQLModel, table=True): # type: ignore
|
||||
default=None,
|
||||
description="Currency of the cashu-token",
|
||||
)
|
||||
parent_key_hash: str | None = Field(
|
||||
default=None, foreign_key="api_keys.hashed_key", index=True
|
||||
)
|
||||
balance_limit: int | None = Field(
|
||||
default=None,
|
||||
description="Max spendable balance in msats for this key (mostly for child keys)",
|
||||
)
|
||||
balance_limit_reset: str | None = Field(
|
||||
default=None,
|
||||
description="Reset policy for balance limit (manual, daily, monthly, etc.)",
|
||||
)
|
||||
balance_limit_reset_date: int | None = Field(
|
||||
default=None,
|
||||
description="Unix timestamp of the last time the balance limit was reset",
|
||||
)
|
||||
validity_date: int | None = Field(
|
||||
default=None,
|
||||
description="Unix timestamp after which the key is no longer valid",
|
||||
)
|
||||
|
||||
@property
|
||||
def total_balance(self) -> int:
|
||||
@@ -54,11 +79,10 @@ class ApiKey(SQLModel, table=True): # type: ignore
|
||||
|
||||
|
||||
async def reset_all_reserved_balances(session: AsyncSession) -> None:
|
||||
logger.info("Resetting all reserved balances to 0")
|
||||
stmt = update(ApiKey).values(reserved_balance=0)
|
||||
await session.exec(stmt) # type: ignore[call-overload]
|
||||
await session.commit()
|
||||
logger.info("Reserved balances reset successfully")
|
||||
logger.info("Reset reserved balances on startup")
|
||||
|
||||
|
||||
class ModelRow(SQLModel, table=True): # type: ignore
|
||||
@@ -81,6 +105,10 @@ class ModelRow(SQLModel, table=True): # type: ignore
|
||||
default=None, description="JSON array of model alias IDs"
|
||||
)
|
||||
enabled: bool = Field(default=True, description="Whether this model is enabled")
|
||||
forwarded_model_id: str | None = Field(
|
||||
default=None,
|
||||
description="Model ID to use when forwarding requests to upstream provider. Defaults to id if not set.",
|
||||
)
|
||||
upstream_provider: "UpstreamProviderRow" = Relationship(back_populates="models")
|
||||
|
||||
|
||||
@@ -106,13 +134,85 @@ class LightningInvoice(SQLModel, table=True): # type: ignore
|
||||
paid_at: int | None = Field(default=None, description="Unix timestamp when paid")
|
||||
|
||||
|
||||
class CashuTransaction(SQLModel, table=True): # type: ignore
|
||||
__tablename__ = "cashu_transactions"
|
||||
|
||||
id: str = Field(
|
||||
primary_key=True,
|
||||
default_factory=lambda: uuid.uuid4().hex,
|
||||
description="Unique transaction identifier",
|
||||
)
|
||||
token: str = Field(description="Serialized Cashu token")
|
||||
amount: int = Field(description="Amount in the token's unit")
|
||||
unit: str = Field(description="Token unit (sat or msat)")
|
||||
mint_url: str | None = Field(default=None, description="Mint URL for the token")
|
||||
type: str = Field(default="out", description="Transaction type: in or out")
|
||||
request_id: str | None = Field(default=None, description="Associated request ID")
|
||||
created_at: int = Field(
|
||||
default_factory=lambda: int(time.time()),
|
||||
description="Unix timestamp",
|
||||
)
|
||||
collected: bool = Field(default=False)
|
||||
swept: bool = Field(default=False)
|
||||
source: str = Field(
|
||||
default="x-cashu",
|
||||
description="Payment source: x-cashu or apikey",
|
||||
)
|
||||
api_key_hashed_key: str | None = Field(
|
||||
default=None,
|
||||
foreign_key="api_keys.hashed_key",
|
||||
index=True,
|
||||
description="Associated API key hash for wallet history",
|
||||
)
|
||||
|
||||
|
||||
async def store_cashu_transaction(
|
||||
token: str,
|
||||
amount: int,
|
||||
unit: str,
|
||||
mint_url: str | None = None,
|
||||
typ: str = "out",
|
||||
request_id: str | None = None,
|
||||
collected: bool = False,
|
||||
created_at: int | None = None,
|
||||
source: str = "x-cashu",
|
||||
api_key_hashed_key: str | None = None,
|
||||
) -> None:
|
||||
try:
|
||||
async with create_session() as session:
|
||||
tx = CashuTransaction(
|
||||
token=token,
|
||||
amount=amount,
|
||||
unit=unit,
|
||||
mint_url=mint_url,
|
||||
type=typ,
|
||||
request_id=request_id,
|
||||
collected=collected,
|
||||
created_at=created_at or int(time.time()),
|
||||
source=source,
|
||||
api_key_hashed_key=api_key_hashed_key,
|
||||
)
|
||||
session.add(tx)
|
||||
await session.commit()
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Failed to store cashu transaction: {e} (type={typ})",
|
||||
extra={"error": str(e), "type": typ},
|
||||
)
|
||||
|
||||
|
||||
class UpstreamProviderRow(SQLModel, table=True): # type: ignore
|
||||
__tablename__ = "upstream_providers"
|
||||
__table_args__ = (
|
||||
UniqueConstraint(
|
||||
"base_url", "api_key", name="uq_upstream_providers_base_url_api_key"
|
||||
),
|
||||
)
|
||||
id: int | None = Field(default=None, primary_key=True)
|
||||
provider_type: str = Field(
|
||||
description="Provider type: custom, openai, anthropic, azure, openrouter, etc."
|
||||
)
|
||||
base_url: str = Field(unique=True, description="Base URL of the upstream API")
|
||||
base_url: str = Field(description="Base URL of the upstream API")
|
||||
api_key: str = Field(description="API key for the upstream provider")
|
||||
api_version: str | None = Field(
|
||||
default=None, description="API version for Azure OpenAI"
|
||||
@@ -121,12 +221,75 @@ class UpstreamProviderRow(SQLModel, table=True): # type: ignore
|
||||
provider_fee: float = Field(
|
||||
default=1.01, description="Provider fee multiplier (default 1%)"
|
||||
)
|
||||
provider_settings: str | None = Field(
|
||||
default=None, description="JSON string for provider-specific settings"
|
||||
)
|
||||
models: list["ModelRow"] = Relationship(
|
||||
back_populates="upstream_provider",
|
||||
sa_relationship_kwargs={"cascade": "all, delete-orphan"},
|
||||
)
|
||||
|
||||
|
||||
class RoutstrFee(SQLModel, table=True): # type: ignore
|
||||
__tablename__ = "routstr_fees"
|
||||
id: int = Field(default=1, primary_key=True)
|
||||
accumulated_msats: int = Field(default=0)
|
||||
total_paid_msats: int = Field(default=0)
|
||||
last_paid_at: int | None = Field(default=None)
|
||||
|
||||
|
||||
class CliToken(SQLModel, table=True): # type: ignore
|
||||
"""Long-lived authorization token for CLI/agent use against admin endpoints."""
|
||||
|
||||
__tablename__ = "cli_tokens"
|
||||
id: str = Field(
|
||||
primary_key=True, default_factory=lambda: uuid.uuid4().hex
|
||||
)
|
||||
token: str = Field(unique=True, index=True, description="Bearer token value")
|
||||
name: str = Field(description="Human-readable label for this token")
|
||||
created_at: int = Field(default_factory=lambda: int(time.time()))
|
||||
last_used_at: int | None = Field(default=None)
|
||||
expires_at: int | None = Field(
|
||||
default=None, description="Optional expiry unix timestamp; null = never expires"
|
||||
)
|
||||
|
||||
|
||||
async def accumulate_routstr_fee(session: AsyncSession, amount_msats: int) -> None:
|
||||
stmt = (
|
||||
update(RoutstrFee)
|
||||
.where(col(RoutstrFee.id) == 1)
|
||||
.values(accumulated_msats=RoutstrFee.accumulated_msats + amount_msats)
|
||||
)
|
||||
result = await session.exec(stmt) # type: ignore[call-overload]
|
||||
if result.rowcount == 0:
|
||||
session.add(RoutstrFee(id=1, accumulated_msats=amount_msats))
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def get_routstr_fee(session: AsyncSession) -> RoutstrFee:
|
||||
fee = await session.get(RoutstrFee, 1)
|
||||
if fee is None:
|
||||
fee = RoutstrFee(id=1, accumulated_msats=0, total_paid_msats=0)
|
||||
session.add(fee)
|
||||
await session.commit()
|
||||
await session.refresh(fee)
|
||||
return fee
|
||||
|
||||
|
||||
async def reset_routstr_fee(session: AsyncSession, paid_msats: int) -> None:
|
||||
stmt = (
|
||||
update(RoutstrFee)
|
||||
.where(col(RoutstrFee.id) == 1)
|
||||
.values(
|
||||
accumulated_msats=RoutstrFee.accumulated_msats - paid_msats,
|
||||
total_paid_msats=RoutstrFee.total_paid_msats + paid_msats,
|
||||
last_paid_at=int(time.time()),
|
||||
)
|
||||
)
|
||||
await session.exec(stmt) # type: ignore[call-overload]
|
||||
await session.commit()
|
||||
|
||||
|
||||
async def balances_for_mint_and_unit(
|
||||
db_session: AsyncSession, mint_url: str, unit: str
|
||||
) -> int:
|
||||
@@ -156,11 +319,64 @@ async def create_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
yield session
|
||||
|
||||
|
||||
def fix_cashu_migrations() -> None:
|
||||
"""
|
||||
Fixes Cashu wallet migrations that are not idempotent.
|
||||
This specifically addresses the 'duplicate column name: public_keys' error
|
||||
in the keysets table of Cashu's internal SQLite databases.
|
||||
"""
|
||||
project_root = pathlib.Path(__file__).resolve().parents[2]
|
||||
wallet_dir = project_root / ".wallet"
|
||||
|
||||
if not wallet_dir.exists() or not wallet_dir.is_dir():
|
||||
return
|
||||
|
||||
logger.info("Checking Cashu wallet databases for migration idempotency")
|
||||
|
||||
for db_file in wallet_dir.glob("*.sqlite3"):
|
||||
try:
|
||||
conn = sqlite3.connect(db_file)
|
||||
cursor = conn.cursor()
|
||||
|
||||
# Check if keysets table exists
|
||||
cursor.execute(
|
||||
"SELECT name FROM sqlite_master WHERE type='table' AND name='keysets'"
|
||||
)
|
||||
if not cursor.fetchone():
|
||||
conn.close()
|
||||
continue
|
||||
|
||||
# Check if public_keys column exists
|
||||
cursor.execute("PRAGMA table_info(keysets)")
|
||||
columns = [info[1] for info in cursor.fetchall()]
|
||||
|
||||
if "public_keys" not in columns:
|
||||
logger.info(f"Adding missing public_keys column to {db_file.name}")
|
||||
cursor.execute("ALTER TABLE keysets ADD COLUMN public_keys TEXT")
|
||||
conn.commit()
|
||||
|
||||
conn.close()
|
||||
except Exception as e:
|
||||
logger.warning(f"Could not check/fix Cashu database {db_file}: {e}")
|
||||
|
||||
|
||||
def _clear_alembic_version() -> None:
|
||||
"""Clear the alembic_version table so stamp/upgrade can proceed."""
|
||||
sync_url = DATABASE_URL.replace("+aiosqlite", "")
|
||||
from sqlalchemy import create_engine, text
|
||||
|
||||
eng = create_engine(sync_url)
|
||||
with eng.begin() as conn:
|
||||
conn.execute(text("DELETE FROM alembic_version"))
|
||||
eng.dispose()
|
||||
|
||||
|
||||
def run_migrations() -> None:
|
||||
"""Run Alembic migrations programmatically."""
|
||||
import pathlib
|
||||
|
||||
try:
|
||||
# Run Cashu migration fix first
|
||||
fix_cashu_migrations()
|
||||
|
||||
# Get the path to the alembic.ini file
|
||||
project_root = pathlib.Path(__file__).resolve().parents[2]
|
||||
alembic_ini_path = project_root / "alembic.ini"
|
||||
@@ -176,8 +392,30 @@ def run_migrations() -> None:
|
||||
# Set the database URL in the config
|
||||
alembic_cfg.set_main_option("sqlalchemy.url", DATABASE_URL)
|
||||
|
||||
# Run migrations to the latest revision
|
||||
command.upgrade(alembic_cfg, "head")
|
||||
try:
|
||||
command.upgrade(alembic_cfg, "head")
|
||||
except CommandError as e:
|
||||
if "Can't locate revision" in str(e):
|
||||
logger.warning(
|
||||
"Database stamped with unknown revision (likely from another branch). "
|
||||
"Re-stamping to current head.",
|
||||
extra={"error": str(e)},
|
||||
)
|
||||
_clear_alembic_version()
|
||||
command.stamp(alembic_cfg, "head")
|
||||
else:
|
||||
raise
|
||||
except OperationalError as e:
|
||||
if "duplicate column name" in str(e).lower():
|
||||
logger.warning(
|
||||
"Migration hit a column that already exists (likely added via "
|
||||
"create_all on another branch). Stamping to current head.",
|
||||
extra={"error": str(e)},
|
||||
)
|
||||
_clear_alembic_version()
|
||||
command.stamp(alembic_cfg, "head")
|
||||
else:
|
||||
raise
|
||||
|
||||
logger.info("Database migrations completed successfully")
|
||||
|
||||
|
||||
@@ -6,6 +6,15 @@ from .logging import get_logger
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class UpstreamError(Exception):
|
||||
"""Exception raised when an upstream provider fails."""
|
||||
|
||||
def __init__(self, message: str, status_code: int = 502):
|
||||
self.message = message
|
||||
self.status_code = status_code
|
||||
super().__init__(message)
|
||||
|
||||
|
||||
async def http_exception_handler(request: Request, exc: Exception) -> JSONResponse:
|
||||
"""Handle HTTP exceptions and include request ID in response."""
|
||||
request_id = getattr(request.state, "request_id", "unknown")
|
||||
@@ -13,16 +22,20 @@ async def http_exception_handler(request: Request, exc: Exception) -> JSONRespon
|
||||
# Get status code and detail - works for both FastAPI and Starlette HTTPException
|
||||
status_code = getattr(exc, "status_code", 500)
|
||||
detail = getattr(exc, "detail", str(exc))
|
||||
path = request.url.path
|
||||
|
||||
logger.warning(
|
||||
"HTTP exception",
|
||||
extra={
|
||||
"request_id": request_id,
|
||||
"status_code": status_code,
|
||||
"detail": detail,
|
||||
"path": request.url.path,
|
||||
},
|
||||
)
|
||||
# 4xx is client behaviour; the uvicorn access log already records it.
|
||||
# Only 5xx warrants a server-side warning/error log here.
|
||||
if status_code >= 500:
|
||||
logger.error(
|
||||
f"HTTP {status_code} on {path}: {detail}",
|
||||
extra={
|
||||
"request_id": request_id,
|
||||
"status_code": status_code,
|
||||
"detail": detail,
|
||||
"path": path,
|
||||
},
|
||||
)
|
||||
|
||||
return JSONResponse(
|
||||
status_code=status_code,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -3,50 +3,61 @@ Logging configuration for Routstr.
|
||||
|
||||
CRITICAL LOG MESSAGES FOR USAGE STATISTICS:
|
||||
===========================================
|
||||
The following log messages are parsed by the usage tracking system (routstr/core/admin.py).
|
||||
The following log messages are parsed by the usage tracking system
|
||||
(routstr/core/usage_analytics_store.py and routstr/core/log_manager.py).
|
||||
DO NOT modify or remove these messages without updating the usage tracking logic:
|
||||
|
||||
1. "Received proxy request" (INFO) - routstr/proxy.py
|
||||
- Used to count total incoming requests
|
||||
- Includes model information in context
|
||||
|
||||
2. "Payment adjustment completed for streaming" (INFO) - routstr/upstream/base.py
|
||||
"Payment adjustment completed for non-streaming" (INFO) - routstr/upstream/base.py
|
||||
2. "Calculated token-based cost" (INFO) - routstr/auth.py
|
||||
- Used to track successful completions and revenue
|
||||
- The 'cost_data.total_msats' field is extracted for revenue calculation
|
||||
- Must include 'cost_data' in extra dict
|
||||
- The 'token_cost', 'model', 'input_tokens', and 'output_tokens' fields are extracted for dashboard metrics
|
||||
|
||||
3. "Payment processed successfully" (INFO) - routstr/auth.py
|
||||
3. "Max cost payment finalized" (INFO) - routstr/auth.py
|
||||
- Used as the successful completion fallback when token usage is unavailable
|
||||
- The 'charged_amount', 'model', 'input_tokens', and 'output_tokens' fields are extracted for dashboard metrics
|
||||
|
||||
4. "Payment processed successfully" (INFO) - routstr/auth.py
|
||||
- Used to count successful payment processing events
|
||||
- Tracks payment-related metrics
|
||||
|
||||
4. "Upstream request failed, revert payment" (WARNING) - routstr/proxy.py
|
||||
5. "Upstream request failed, revert payment" (WARNING) - routstr/proxy.py
|
||||
- Used to track failed requests and refunds
|
||||
- The 'max_cost_for_model' field is extracted for refund calculation
|
||||
- Must include 'max_cost_for_model' in extra dict
|
||||
|
||||
5. Any ERROR level logs with "upstream" in the message
|
||||
6. Any ERROR level logs with "upstream" in the message
|
||||
- Used to count upstream provider errors
|
||||
- Helps identify service reliability issues
|
||||
|
||||
If you need to modify these messages, ensure you also update the parsing logic in:
|
||||
- routstr/core/admin.py:_aggregate_metrics_by_time()
|
||||
- routstr/core/admin.py:_get_summary_stats()
|
||||
- routstr/core/admin.py:get_revenue_by_model()
|
||||
- routstr/core/usage_analytics_store.py
|
||||
- routstr/core/log_manager.py
|
||||
"""
|
||||
|
||||
import logging.config
|
||||
import logging.handlers
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
import tomllib
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from pythonjsonlogger import jsonlogger
|
||||
from rich.console import Console
|
||||
from rich.logging import RichHandler
|
||||
|
||||
# Only use RichHandler when stdout is a real TTY. In non-TTY contexts
|
||||
# (docker logs, pipes, CI) Rich pads every line to width and wraps long
|
||||
# records, producing visually-empty trailing whitespace and split records.
|
||||
# A plain StreamHandler avoids both problems.
|
||||
_stdout_is_tty = sys.stdout.isatty()
|
||||
_console = Console(soft_wrap=True) if _stdout_is_tty else None
|
||||
|
||||
# Define custom TRACE level
|
||||
TRACE_LEVEL = 5
|
||||
logging.addLevelName(TRACE_LEVEL, "TRACE")
|
||||
@@ -259,6 +270,26 @@ def setup_logging() -> None:
|
||||
if console_enabled:
|
||||
handlers.append("console")
|
||||
|
||||
if _stdout_is_tty:
|
||||
console_handler: dict[str, Any] = {
|
||||
"()": RichHandler,
|
||||
"level": log_level,
|
||||
"show_time": False,
|
||||
"show_path": False,
|
||||
"rich_tracebacks": True,
|
||||
"markup": True,
|
||||
"console": _console,
|
||||
"filters": ["request_id_filter", "security_filter"],
|
||||
}
|
||||
else:
|
||||
console_handler = {
|
||||
"class": "logging.StreamHandler",
|
||||
"level": log_level,
|
||||
"formatter": "plain",
|
||||
"stream": "ext://sys.stdout",
|
||||
"filters": ["request_id_filter", "security_filter"],
|
||||
}
|
||||
|
||||
LOGGING_CONFIG = {
|
||||
"version": 1,
|
||||
"disable_existing_loggers": False,
|
||||
@@ -268,6 +299,10 @@ def setup_logging() -> None:
|
||||
"format": "%(asctime)s %(name)s %(levelname)s %(message)s %(pathname)s %(lineno)d %(version)s %(request_id)s",
|
||||
"datefmt": "%Y-%m-%d %H:%M:%S",
|
||||
},
|
||||
"plain": {
|
||||
"format": "%(asctime)s %(levelname)-7s %(name)s %(message)s",
|
||||
"datefmt": "%Y-%m-%d %H:%M:%S",
|
||||
},
|
||||
},
|
||||
"filters": {
|
||||
"version_filter": {"()": VersionFilter},
|
||||
@@ -275,15 +310,7 @@ def setup_logging() -> None:
|
||||
"security_filter": {"()": SecurityFilter},
|
||||
},
|
||||
"handlers": {
|
||||
"console": {
|
||||
"()": RichHandler,
|
||||
"level": log_level,
|
||||
"show_time": False,
|
||||
"show_path": False,
|
||||
"rich_tracebacks": True,
|
||||
"markup": True,
|
||||
"filters": ["request_id_filter", "security_filter"],
|
||||
},
|
||||
"console": console_handler,
|
||||
"file": {
|
||||
"()": DailyRotatingFileHandler,
|
||||
"level": log_level,
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
import asyncio
|
||||
import os
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
from typing import AsyncGenerator
|
||||
@@ -9,34 +8,38 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import FileResponse, RedirectResponse
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
from starlette.exceptions import HTTPException
|
||||
from starlette.responses import Response as StarletteResponse
|
||||
from starlette.types import Scope
|
||||
|
||||
from ..auth import periodic_key_reset
|
||||
from ..balance import balance_router, deprecated_wallet_router
|
||||
from ..discovery import providers_cache_refresher, providers_router
|
||||
from ..nip91 import announce_provider
|
||||
from ..payment.models import (
|
||||
models_router,
|
||||
update_sats_pricing,
|
||||
from ..lightning import lightning_router, periodic_invoice_watcher
|
||||
from ..nostr import (
|
||||
announce_provider,
|
||||
providers_cache_refresher,
|
||||
publish_usage_analytics,
|
||||
)
|
||||
from ..nostr.discovery import providers_router
|
||||
from ..payment.models import models_router, update_sats_pricing
|
||||
from ..payment.price import update_prices_periodically
|
||||
from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically
|
||||
from ..wallet import periodic_payout
|
||||
from ..upstream.auto_topup import periodic_auto_topup
|
||||
from ..upstream.litellm_routing import configure_litellm
|
||||
from ..wallet import periodic_payout, periodic_refund_sweep, periodic_routstr_fee_payout
|
||||
from .admin import admin_router
|
||||
from .db import create_session, init_db, run_migrations
|
||||
from .exceptions import general_exception_handler, http_exception_handler
|
||||
from .logging import get_logger, setup_logging
|
||||
from .middleware import LoggingMiddleware
|
||||
from .not_found import _NOT_FOUND_HTML, not_found_catch_all # noqa: F401
|
||||
from .settings import SettingsService
|
||||
from .settings import settings as global_settings
|
||||
from .version import __version__
|
||||
|
||||
# Initialize logging first
|
||||
setup_logging()
|
||||
logger = get_logger(__name__)
|
||||
|
||||
if os.getenv("VERSION_SUFFIX") is not None:
|
||||
__version__ = f"0.2.2-{os.getenv('VERSION_SUFFIX')}"
|
||||
else:
|
||||
__version__ = "0.2.2"
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
@@ -46,11 +49,21 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
pricing_task = None
|
||||
payout_task = None
|
||||
nip91_task = None
|
||||
analytics_task = None
|
||||
providers_task = None
|
||||
models_refresh_task = None
|
||||
model_maps_refresh_task = None
|
||||
key_reset_task = None
|
||||
auto_topup_task = None
|
||||
refund_sweep_task = None
|
||||
routstr_fee_task = None
|
||||
invoice_watcher_task = None
|
||||
|
||||
try:
|
||||
# Apply litellm-wide settings (drop_params, chat-completions URL,
|
||||
# debug logging) before any upstream provider dispatches a request.
|
||||
configure_litellm()
|
||||
|
||||
# Run database migrations on startup
|
||||
run_migrations()
|
||||
|
||||
@@ -95,15 +108,24 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
btc_price_task = asyncio.create_task(update_prices_periodically())
|
||||
pricing_task = asyncio.create_task(update_sats_pricing())
|
||||
if global_settings.models_refresh_interval_seconds > 0:
|
||||
# Pass the accessor (not its current value) so the loop sees providers
|
||||
# added/changed via reinitialize_upstreams() instead of staying pinned
|
||||
# to the startup snapshot.
|
||||
models_refresh_task = asyncio.create_task(
|
||||
refresh_upstreams_models_periodically(get_upstreams())
|
||||
refresh_upstreams_models_periodically(get_upstreams)
|
||||
)
|
||||
model_maps_refresh_task = asyncio.create_task(refresh_model_maps_periodically())
|
||||
payout_task = asyncio.create_task(periodic_payout())
|
||||
if global_settings.nsec:
|
||||
nip91_task = asyncio.create_task(announce_provider())
|
||||
analytics_task = asyncio.create_task(publish_usage_analytics())
|
||||
if global_settings.providers_refresh_interval_seconds > 0:
|
||||
providers_task = asyncio.create_task(providers_cache_refresher())
|
||||
key_reset_task = asyncio.create_task(periodic_key_reset())
|
||||
auto_topup_task = asyncio.create_task(periodic_auto_topup())
|
||||
refund_sweep_task = asyncio.create_task(periodic_refund_sweep())
|
||||
routstr_fee_task = asyncio.create_task(periodic_routstr_fee_payout())
|
||||
invoice_watcher_task = asyncio.create_task(periodic_invoice_watcher())
|
||||
|
||||
yield
|
||||
|
||||
@@ -127,12 +149,24 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
payout_task.cancel()
|
||||
if nip91_task is not None:
|
||||
nip91_task.cancel()
|
||||
if analytics_task is not None:
|
||||
analytics_task.cancel()
|
||||
if providers_task is not None:
|
||||
providers_task.cancel()
|
||||
if models_refresh_task is not None:
|
||||
models_refresh_task.cancel()
|
||||
if model_maps_refresh_task is not None:
|
||||
model_maps_refresh_task.cancel()
|
||||
if key_reset_task is not None:
|
||||
key_reset_task.cancel()
|
||||
if auto_topup_task is not None:
|
||||
auto_topup_task.cancel()
|
||||
if refund_sweep_task is not None:
|
||||
refund_sweep_task.cancel()
|
||||
if routstr_fee_task is not None:
|
||||
routstr_fee_task.cancel()
|
||||
if invoice_watcher_task is not None:
|
||||
invoice_watcher_task.cancel()
|
||||
|
||||
try:
|
||||
tasks_to_wait = []
|
||||
@@ -144,12 +178,24 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
tasks_to_wait.append(payout_task)
|
||||
if nip91_task is not None:
|
||||
tasks_to_wait.append(nip91_task)
|
||||
if analytics_task is not None:
|
||||
tasks_to_wait.append(analytics_task)
|
||||
if providers_task is not None:
|
||||
tasks_to_wait.append(providers_task)
|
||||
if models_refresh_task is not None:
|
||||
tasks_to_wait.append(models_refresh_task)
|
||||
if model_maps_refresh_task is not None:
|
||||
tasks_to_wait.append(model_maps_refresh_task)
|
||||
if key_reset_task is not None:
|
||||
tasks_to_wait.append(key_reset_task)
|
||||
if auto_topup_task is not None:
|
||||
tasks_to_wait.append(auto_topup_task)
|
||||
if refund_sweep_task is not None:
|
||||
tasks_to_wait.append(refund_sweep_task)
|
||||
if routstr_fee_task is not None:
|
||||
tasks_to_wait.append(routstr_fee_task)
|
||||
if invoice_watcher_task is not None:
|
||||
tasks_to_wait.append(invoice_watcher_task)
|
||||
|
||||
if tasks_to_wait:
|
||||
await asyncio.gather(*tasks_to_wait, return_exceptions=True)
|
||||
@@ -161,6 +207,23 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
|
||||
)
|
||||
|
||||
|
||||
class _ImmutableStaticFiles(StaticFiles):
|
||||
"""Static files with long Cache-Control for content-hashed Next.js assets.
|
||||
|
||||
Files under `/_next/static/` are emitted with content hashes in their
|
||||
filenames and never mutate, so we serve them with a one-year immutable
|
||||
cache header so browsers and CDNs stop revalidating on every reload.
|
||||
"""
|
||||
|
||||
async def get_response(self, path: str, scope: Scope) -> StarletteResponse:
|
||||
response = await super().get_response(path, scope)
|
||||
if response.status_code == 200:
|
||||
response.headers["Cache-Control"] = (
|
||||
"public, max-age=31536000, immutable"
|
||||
)
|
||||
return response
|
||||
|
||||
|
||||
app = FastAPI(version=__version__, lifespan=lifespan)
|
||||
|
||||
|
||||
@@ -170,7 +233,7 @@ app.add_middleware(
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
expose_headers=["x-routstr-request-id"],
|
||||
expose_headers=["x-routstr-request-id", "x-cashu"],
|
||||
)
|
||||
|
||||
# Add logging middleware
|
||||
@@ -191,6 +254,7 @@ async def info() -> dict:
|
||||
"mints": global_settings.cashu_mints,
|
||||
"http_url": global_settings.http_url,
|
||||
"onion_url": global_settings.onion_url,
|
||||
"child_key_cost_msats": global_settings.child_key_cost,
|
||||
}
|
||||
|
||||
|
||||
@@ -206,7 +270,7 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir():
|
||||
|
||||
app.mount(
|
||||
"/_next",
|
||||
StaticFiles(directory=UI_DIST_PATH / "_next", check_dir=True),
|
||||
_ImmutableStaticFiles(directory=UI_DIST_PATH / "_next", check_dir=True),
|
||||
name="next-static",
|
||||
)
|
||||
|
||||
@@ -214,100 +278,70 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir():
|
||||
async def serve_root_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "index.html")
|
||||
|
||||
# Add explicit route for /index.txt to redirect to /
|
||||
# Serve the App Router RSC payload for the home page.
|
||||
@app.get("/index.txt", include_in_schema=False)
|
||||
async def redirect_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/")
|
||||
async def serve_root_rsc() -> FileResponse:
|
||||
return FileResponse(
|
||||
UI_DIST_PATH / "index.txt", media_type="text/x-component"
|
||||
)
|
||||
|
||||
# Next.js is built with `trailingSlash: true`, so all UI page URLs end
|
||||
# with a slash (e.g. `/login/`). The proxy router catches `/{path:path}`
|
||||
# before FastAPI's `redirect_slashes` logic can normalize the URL, so we
|
||||
# must register both the with-slash and without-slash variants here.
|
||||
UI_PAGES = (
|
||||
"dashboard",
|
||||
"login",
|
||||
"model",
|
||||
"providers",
|
||||
"settings",
|
||||
"transactions",
|
||||
"balances",
|
||||
"logs",
|
||||
"usage",
|
||||
"unauthorized",
|
||||
)
|
||||
|
||||
def _register_ui_page(name: str) -> None:
|
||||
page_dir = UI_DIST_PATH / name
|
||||
index_html = page_dir / "index.html"
|
||||
index_txt = page_dir / "index.txt"
|
||||
|
||||
async def serve_page() -> FileResponse:
|
||||
return FileResponse(index_html)
|
||||
|
||||
async def serve_page_rsc() -> FileResponse:
|
||||
return FileResponse(index_txt, media_type="text/x-component")
|
||||
|
||||
app.add_api_route(
|
||||
f"/{name}",
|
||||
serve_page,
|
||||
methods=["GET"],
|
||||
include_in_schema=False,
|
||||
name=f"serve_{name}_ui",
|
||||
)
|
||||
app.add_api_route(
|
||||
f"/{name}/",
|
||||
serve_page,
|
||||
methods=["GET"],
|
||||
include_in_schema=False,
|
||||
name=f"serve_{name}_ui_slash",
|
||||
)
|
||||
app.add_api_route(
|
||||
f"/{name}/index.txt",
|
||||
serve_page_rsc,
|
||||
methods=["GET"],
|
||||
include_in_schema=False,
|
||||
name=f"serve_{name}_rsc",
|
||||
)
|
||||
|
||||
for _page in UI_PAGES:
|
||||
_register_ui_page(_page)
|
||||
|
||||
@app.get("/admin")
|
||||
async def admin_redirect() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "index.html")
|
||||
|
||||
@app.get("/dashboard", include_in_schema=False)
|
||||
async def serve_dashboard_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "index.html")
|
||||
|
||||
@app.get("/login", include_in_schema=False)
|
||||
async def serve_login_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "login" / "index.html")
|
||||
|
||||
# Add explicit route for /login/index.txt to redirect to /login
|
||||
@app.get("/login/index.txt", include_in_schema=False)
|
||||
async def redirect_login_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/login")
|
||||
|
||||
@app.get("/model", include_in_schema=False)
|
||||
async def serve_models_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "model" / "index.html")
|
||||
|
||||
# Add explicit route for /model/index.txt to redirect to /model
|
||||
@app.get("/model/index.txt", include_in_schema=False)
|
||||
async def redirect_model_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/model")
|
||||
|
||||
@app.get("/providers", include_in_schema=False)
|
||||
async def serve_providers_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "providers" / "index.html")
|
||||
|
||||
# Add explicit route for /providers/index.txt to redirect to /providers
|
||||
@app.get("/providers/index.txt", include_in_schema=False)
|
||||
async def redirect_providers_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/providers")
|
||||
|
||||
@app.get("/settings", include_in_schema=False)
|
||||
async def serve_settings_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "settings" / "index.html")
|
||||
|
||||
# Add explicit route for /settings/index.txt to redirect to /settings
|
||||
@app.get("/settings/index.txt", include_in_schema=False)
|
||||
async def redirect_settings_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/settings")
|
||||
|
||||
@app.get("/transactions", include_in_schema=False)
|
||||
async def serve_transactions_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "transactions" / "index.html")
|
||||
|
||||
# Add explicit route for /transactions/index.txt to redirect to /transactions
|
||||
@app.get("/transactions/index.txt", include_in_schema=False)
|
||||
async def redirect_transactions_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/transactions")
|
||||
|
||||
@app.get("/balances", include_in_schema=False)
|
||||
async def serve_balances_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "balances" / "index.html")
|
||||
|
||||
# Add explicit route for /balances/index.txt to redirect to /balances
|
||||
@app.get("/balances/index.txt", include_in_schema=False)
|
||||
async def redirect_balances_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/balances")
|
||||
|
||||
@app.get("/logs", include_in_schema=False)
|
||||
async def serve_logs_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "logs" / "index.html")
|
||||
|
||||
# Add explicit route for /logs/index.txt to redirect to /logs
|
||||
@app.get("/logs/index.txt", include_in_schema=False)
|
||||
async def redirect_logs_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/logs")
|
||||
|
||||
@app.get("/usage", include_in_schema=False)
|
||||
async def serve_usage_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "usage" / "index.html")
|
||||
|
||||
# Add explicit route for /usage/index.txt to redirect to /usage
|
||||
@app.get("/usage/index.txt", include_in_schema=False)
|
||||
async def redirect_usage_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/usage")
|
||||
|
||||
@app.get("/unauthorized", include_in_schema=False)
|
||||
async def serve_unauthorized_ui() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "unauthorized" / "index.html")
|
||||
|
||||
# Add explicit route for /unauthorized/index.txt to redirect to /unauthorized
|
||||
@app.get("/unauthorized/index.txt", include_in_schema=False)
|
||||
async def redirect_unauthorized_index_txt() -> RedirectResponse:
|
||||
return RedirectResponse("/unauthorized")
|
||||
|
||||
@app.get("/favicon.ico", include_in_schema=False)
|
||||
async def serve_favicon() -> FileResponse:
|
||||
icon_path = UI_DIST_PATH / "icon.ico"
|
||||
@@ -319,9 +353,6 @@ if UI_DIST_PATH.exists() and UI_DIST_PATH.is_dir():
|
||||
async def serve_icon() -> FileResponse:
|
||||
return FileResponse(UI_DIST_PATH / "icon.ico")
|
||||
|
||||
app.mount(
|
||||
"/static", StaticFiles(directory=UI_DIST_PATH, check_dir=True), name="ui-static"
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
f"UI dist directory not found at {UI_DIST_PATH}, skipping static file serving"
|
||||
@@ -341,6 +372,7 @@ else:
|
||||
app.include_router(models_router)
|
||||
app.include_router(admin_router)
|
||||
app.include_router(balance_router)
|
||||
app.include_router(lightning_router)
|
||||
app.include_router(deprecated_wallet_router)
|
||||
app.include_router(providers_router)
|
||||
app.include_router(proxy_router)
|
||||
|
||||
@@ -14,8 +14,54 @@ logger = get_logger(__name__)
|
||||
request_id_context: ContextVar[str | None] = ContextVar("request_id")
|
||||
|
||||
|
||||
# Methods that are never logged: HEAD requests are health probes from
|
||||
# monitoring/load balancers, OPTIONS are CORS preflights — both are framework
|
||||
# chatter, not user-meaningful events.
|
||||
_SKIP_LOG_METHODS: frozenset[str] = frozenset({"HEAD", "OPTIONS"})
|
||||
|
||||
# Path prefixes to skip. Includes Next.js static chunks and the admin
|
||||
# dashboard's internal polling API (/admin/api/*) which the UI hits on a timer
|
||||
# to refresh balances, logs, providers, etc. — high volume, low diagnostic
|
||||
# value. Mutating admin actions are recorded separately in the audit log.
|
||||
_SKIP_LOG_PREFIXES: tuple[str, ...] = (
|
||||
"/_next/",
|
||||
"/admin/api/",
|
||||
)
|
||||
|
||||
# Exact paths to skip. RSC payload prefetches (`*/index.txt`) fire automatically
|
||||
# as the user hovers near `<Link>`s, and `/v1/wallet/info` is polled by the UI.
|
||||
_SKIP_LOG_EXACT: frozenset[str] = frozenset(
|
||||
{
|
||||
"/favicon.ico",
|
||||
"/icon.ico",
|
||||
"/v1/wallet/info",
|
||||
"/index.txt",
|
||||
"/login/index.txt",
|
||||
"/model/index.txt",
|
||||
"/providers/index.txt",
|
||||
"/settings/index.txt",
|
||||
"/transactions/index.txt",
|
||||
"/balances/index.txt",
|
||||
"/logs/index.txt",
|
||||
"/usage/index.txt",
|
||||
"/unauthorized/index.txt",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _should_log(method: str, path: str) -> bool:
|
||||
if method in _SKIP_LOG_METHODS:
|
||||
return False
|
||||
if path in _SKIP_LOG_EXACT:
|
||||
return False
|
||||
return not any(path.startswith(prefix) for prefix in _SKIP_LOG_PREFIXES)
|
||||
|
||||
|
||||
class LoggingMiddleware(BaseHTTPMiddleware):
|
||||
"""Middleware to log detailed request and response information."""
|
||||
"""Middleware to log proxy interactions and page navigation.
|
||||
|
||||
Skips logging for static assets and Next.js chunks to avoid noise.
|
||||
"""
|
||||
|
||||
async def dispatch(self, request: Request, call_next: Callable) -> Response:
|
||||
# Generate request ID
|
||||
@@ -25,62 +71,20 @@ class LoggingMiddleware(BaseHTTPMiddleware):
|
||||
# Set request ID in context for logging
|
||||
token = request_id_context.set(request_id)
|
||||
|
||||
path = request.url.path
|
||||
should_log = _should_log(request.method, path)
|
||||
|
||||
# Start timing
|
||||
start_time = time.time()
|
||||
|
||||
# Log request details
|
||||
request_body = None
|
||||
if request.method in ["POST", "PUT", "PATCH"]:
|
||||
try:
|
||||
# Only read body for non-streaming requests
|
||||
if hasattr(request, "_body"):
|
||||
request_body = await request.body()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Extract request info
|
||||
client_host = None
|
||||
if request.client:
|
||||
client_host = request.client.host
|
||||
|
||||
# Log incoming request
|
||||
logger.info(
|
||||
"Incoming request",
|
||||
extra={
|
||||
"request_id": request_id,
|
||||
"method": request.method,
|
||||
"path": request.url.path,
|
||||
"query_params": dict(request.query_params),
|
||||
"client_host": client_host,
|
||||
"headers": {
|
||||
k: v
|
||||
for k, v in request.headers.items()
|
||||
if k.lower()
|
||||
not in [
|
||||
"authorization",
|
||||
"x-cashu",
|
||||
"cookie",
|
||||
"cf-connecting-ip",
|
||||
"cf-ipcountry",
|
||||
"x-forwarded-for",
|
||||
"x-real-ip",
|
||||
]
|
||||
},
|
||||
"body_size": len(request_body) if request_body else 0,
|
||||
},
|
||||
)
|
||||
|
||||
# Log at TRACE level for full body (security filter will redact sensitive data)
|
||||
if request_body and hasattr(logger, "exception"):
|
||||
logger.exception(
|
||||
"Request body",
|
||||
if should_log:
|
||||
logger.info(
|
||||
"Incoming request",
|
||||
extra={
|
||||
"request_id": request_id,
|
||||
"method": request.method,
|
||||
"path": request.url.path,
|
||||
"body": request_body.decode("utf-8", errors="ignore")[
|
||||
:1000
|
||||
], # Limit size
|
||||
"path": path,
|
||||
"query_params": dict(request.query_params),
|
||||
},
|
||||
)
|
||||
|
||||
@@ -88,39 +92,33 @@ class LoggingMiddleware(BaseHTTPMiddleware):
|
||||
try:
|
||||
response = await call_next(request)
|
||||
|
||||
# Calculate duration
|
||||
duration = time.time() - start_time
|
||||
|
||||
# Log response
|
||||
logger.info(
|
||||
"Request completed",
|
||||
extra={
|
||||
"request_id": request_id,
|
||||
"method": request.method,
|
||||
"path": request.url.path,
|
||||
"status_code": response.status_code,
|
||||
"duration_ms": round(duration * 1000, 2),
|
||||
"client_host": client_host,
|
||||
},
|
||||
)
|
||||
if should_log:
|
||||
duration = time.time() - start_time
|
||||
logger.info(
|
||||
"Request completed",
|
||||
extra={
|
||||
"request_id": request_id,
|
||||
"method": request.method,
|
||||
"path": path,
|
||||
"status_code": response.status_code,
|
||||
"duration_ms": round(duration * 1000, 2),
|
||||
},
|
||||
)
|
||||
if hasattr(response, "headers"):
|
||||
response.headers["x-routstr-request-id"] = request_id
|
||||
|
||||
return response
|
||||
|
||||
except Exception as e:
|
||||
# Calculate duration
|
||||
# Always log failures, even for skipped paths, so we don't lose errors.
|
||||
duration = time.time() - start_time
|
||||
|
||||
# Log error
|
||||
logger.error(
|
||||
"Request failed",
|
||||
extra={
|
||||
"request_id": request_id,
|
||||
"method": request.method,
|
||||
"path": request.url.path,
|
||||
"path": path,
|
||||
"duration_ms": round(duration * 1000, 2),
|
||||
"client_host": client_host,
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
},
|
||||
|
||||
55
routstr/core/not_found.py
Normal file
55
routstr/core/not_found.py
Normal file
@@ -0,0 +1,55 @@
|
||||
"""Shared 404 handler used by the proxy catch-all and tests."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from fastapi import Request
|
||||
from fastapi.responses import HTMLResponse, JSONResponse, Response
|
||||
|
||||
_NOT_FOUND_HTML_FILE = Path(__file__).parent.parent.parent / "ui_out" / "404.html"
|
||||
|
||||
|
||||
def _read_not_found_html() -> str | None:
|
||||
try:
|
||||
return _NOT_FOUND_HTML_FILE.read_text(encoding="utf-8")
|
||||
except OSError:
|
||||
return None
|
||||
|
||||
|
||||
_NOT_FOUND_HTML: str | None = _read_not_found_html()
|
||||
|
||||
|
||||
def build_not_found_response(request: Request, path: str) -> Response:
|
||||
"""Return a 404 response.
|
||||
|
||||
HTML 404 page only for GET requests from browsers (Accept: text/html).
|
||||
All POST requests and API clients receive a JSON 404.
|
||||
"""
|
||||
accept = request.headers.get("accept", "").lower()
|
||||
prefers_html = (
|
||||
request.method == "GET"
|
||||
and "text/html" in accept
|
||||
and "application/json" not in accept
|
||||
)
|
||||
request_id = getattr(request.state, "request_id", "unknown")
|
||||
|
||||
if prefers_html and _NOT_FOUND_HTML is not None:
|
||||
return HTMLResponse(content=_NOT_FOUND_HTML, status_code=404)
|
||||
|
||||
return JSONResponse(
|
||||
status_code=404,
|
||||
content={
|
||||
"error": {
|
||||
"message": f"Path '/{path}' not found",
|
||||
"type": "not_found",
|
||||
"code": 404,
|
||||
},
|
||||
"request_id": request_id,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def not_found_catch_all(request: Request, path: str) -> Response:
|
||||
"""ASGI handler form of :func:`build_not_found_response`."""
|
||||
return build_not_found_response(request, path)
|
||||
@@ -41,6 +41,15 @@ class Settings(BaseSettings):
|
||||
primary_mint: str = Field(default="", env="PRIMARY_MINT_URL")
|
||||
primary_mint_unit: str = Field(default="sat", env="PRIMARY_MINT_UNIT")
|
||||
|
||||
# Lightning payout configuration
|
||||
# Minimum available balance (in satoshis) before profit is paid out over
|
||||
# Lightning
|
||||
min_payout_sat: int = Field(default=210, gt=0, env="MIN_PAYOUT_SAT")
|
||||
# Interval (seconds) between periodic payout attempts. Must be positive.
|
||||
payout_interval_seconds: int = Field(
|
||||
default=900, gt=0, env="PAYOUT_INTERVAL_SECONDS"
|
||||
)
|
||||
|
||||
# Pricing
|
||||
# Default behavior: derive pricing from MODELS
|
||||
# If fixed_pricing is True -> use fixed_cost_per_request and ignore tokens
|
||||
@@ -52,6 +61,7 @@ class Settings(BaseSettings):
|
||||
exchange_fee: float = Field(default=1.005, env="EXCHANGE_FEE")
|
||||
upstream_provider_fee: float = Field(default=1.05, env="UPSTREAM_PROVIDER_FEE")
|
||||
tolerance_percentage: float = Field(default=1.0, env="TOLERANCE_PERCENTAGE")
|
||||
child_key_cost: int = Field(default=0, env="CHILD_KEY_COST")
|
||||
# Minimum per-request charge in millisatoshis when model pricing is free/zero
|
||||
min_request_msat: int = Field(default=1, env="MIN_REQUEST_MSAT")
|
||||
reset_reserved_balance_on_startup: bool = Field(
|
||||
@@ -73,6 +83,7 @@ class Settings(BaseSettings):
|
||||
enable_pricing_refresh: bool = Field(default=True, env="ENABLE_PRICING_REFRESH")
|
||||
enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH")
|
||||
refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS")
|
||||
refund_sweep_ttl_seconds: int = Field(default=604800, env="REFUND_SWEEP_TTL_SECONDS")
|
||||
|
||||
# Logging
|
||||
log_level: str = Field(default="INFO", env="LOG_LEVEL")
|
||||
@@ -91,6 +102,20 @@ class Settings(BaseSettings):
|
||||
|
||||
# Discovery
|
||||
relays: list[str] = Field(default_factory=list, env="RELAYS")
|
||||
enable_analytics_sharing: bool = Field(
|
||||
default=True, env="ENABLE_ANALYTICS_SHARING"
|
||||
)
|
||||
|
||||
def _normalize_settings_data(data: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Discard unknown keys from persisted settings."""
|
||||
normalized: dict[str, Any] = {}
|
||||
known_fields = Settings.__fields__
|
||||
|
||||
for key, value in data.items():
|
||||
if key in known_fields:
|
||||
normalized[key] = value
|
||||
|
||||
return normalized
|
||||
|
||||
|
||||
def _compute_primary_mint(cashu_mints: list[str]) -> str:
|
||||
@@ -142,7 +167,7 @@ def resolve_bootstrap() -> Settings:
|
||||
pass
|
||||
if not base.onion_url:
|
||||
try:
|
||||
from ..nip91 import discover_onion_url_from_tor # type: ignore
|
||||
from ..nostr.listing import discover_onion_url_from_tor # type: ignore
|
||||
|
||||
discovered = discover_onion_url_from_tor()
|
||||
if discovered:
|
||||
@@ -230,16 +255,21 @@ class SettingsService:
|
||||
|
||||
db_id, db_data, _updated_at = row
|
||||
try:
|
||||
db_json = (
|
||||
db_json_raw = (
|
||||
json.loads(db_data) if isinstance(db_data, str) else dict(db_data)
|
||||
)
|
||||
if not isinstance(db_json_raw, dict):
|
||||
db_json_raw = {}
|
||||
except Exception:
|
||||
db_json = {}
|
||||
db_json_raw = {}
|
||||
db_json = _normalize_settings_data(db_json_raw)
|
||||
|
||||
valid_fields = set(env_resolved.dict().keys())
|
||||
merged_dict: dict[str, Any] = dict(env_resolved.dict())
|
||||
merged_dict.update(
|
||||
{k: v for k, v in db_json.items() if v not in (None, "", [], {})}
|
||||
{k: v for k, v in db_json.items() if v not in (None, "", [], {}) and k in valid_fields}
|
||||
)
|
||||
merged_dict = Settings(**merged_dict).dict()
|
||||
|
||||
# Ensure primary_mint is consistent with cashu_mints if not explicitly set
|
||||
if not merged_dict.get("primary_mint"):
|
||||
@@ -247,7 +277,7 @@ class SettingsService:
|
||||
merged_dict.get("cashu_mints", [])
|
||||
)
|
||||
|
||||
if any(k not in db_json for k in merged_dict.keys()):
|
||||
if db_json_raw != merged_dict:
|
||||
await db_session.exec( # type: ignore
|
||||
text(
|
||||
"UPDATE settings SET data = :data, updated_at = :updated_at WHERE id = 1"
|
||||
@@ -270,7 +300,7 @@ class SettingsService:
|
||||
) -> Settings:
|
||||
async with cls._lock:
|
||||
current = cls.get()
|
||||
candidate_dict = {**current.dict(), **partial}
|
||||
candidate_dict = {**current.dict(), **_normalize_settings_data(partial)}
|
||||
candidate = Settings(**candidate_dict)
|
||||
from sqlmodel import text
|
||||
|
||||
@@ -304,8 +334,10 @@ class SettingsService:
|
||||
raise RuntimeError("Settings row missing")
|
||||
(data_str,) = row
|
||||
data = json.loads(data_str) if isinstance(data_str, str) else dict(data_str)
|
||||
valid_fields = set(settings.dict().keys())
|
||||
# Update in-place
|
||||
for k, v in data.items():
|
||||
setattr(settings, k, v)
|
||||
if k in valid_fields:
|
||||
setattr(settings, k, v)
|
||||
cls._current = settings
|
||||
return settings
|
||||
|
||||
1380
routstr/core/usage_analytics_store.py
Normal file
1380
routstr/core/usage_analytics_store.py
Normal file
File diff suppressed because it is too large
Load Diff
81
routstr/core/version.py
Normal file
81
routstr/core/version.py
Normal file
@@ -0,0 +1,81 @@
|
||||
"""Application version resolution.
|
||||
|
||||
Priority order:
|
||||
|
||||
1. ``VERSION_SUFFIX`` env var (manual override; preserves prior behaviour).
|
||||
2. Bare base version when HEAD is on the matching release tag (detected via
|
||||
``GIT_TAG`` env or ``git describe --tags --exact-match HEAD``).
|
||||
3. ``GIT_COMMIT`` env var (build-time injection) -> ``<base>+g<sha>``.
|
||||
4. Local ``.git`` lookup (source checkouts) -> ``<base>+g<sha>``.
|
||||
5. Fallback: bare base version.
|
||||
|
||||
The ``+g<sha>`` form is PEP 440 local-version syntax so the result remains a
|
||||
valid package version.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
from functools import lru_cache
|
||||
from pathlib import Path
|
||||
|
||||
BASE_VERSION = "0.4.3"
|
||||
|
||||
_REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
_GIT_TIMEOUT_SECONDS = 2.0
|
||||
|
||||
|
||||
def _run_git(*args: str) -> str | None:
|
||||
try:
|
||||
result = subprocess.run( # noqa: S603 - fixed argv, no shell
|
||||
["git", *args],
|
||||
cwd=_REPO_ROOT,
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=_GIT_TIMEOUT_SECONDS,
|
||||
)
|
||||
except (FileNotFoundError, subprocess.SubprocessError, OSError):
|
||||
return None
|
||||
if result.returncode != 0:
|
||||
return None
|
||||
return result.stdout.strip() or None
|
||||
|
||||
|
||||
def _git_short_sha() -> str | None:
|
||||
sha = os.getenv("GIT_COMMIT", "").strip()
|
||||
if sha:
|
||||
return sha[:7]
|
||||
return _run_git("rev-parse", "--short=7", "HEAD")
|
||||
|
||||
|
||||
def _on_tagged_release() -> bool:
|
||||
tag = os.getenv("GIT_TAG", "").strip()
|
||||
if tag:
|
||||
return tag.lstrip("v") == BASE_VERSION
|
||||
described = _run_git("describe", "--tags", "--exact-match", "HEAD")
|
||||
if described and described.lstrip("v") == BASE_VERSION:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def get_version() -> str:
|
||||
suffix = os.getenv("VERSION_SUFFIX")
|
||||
if suffix is not None:
|
||||
return f"{BASE_VERSION}-{suffix}"
|
||||
|
||||
if _on_tagged_release():
|
||||
return BASE_VERSION
|
||||
|
||||
sha = _git_short_sha()
|
||||
if not sha:
|
||||
return BASE_VERSION
|
||||
|
||||
return f"{BASE_VERSION}+g{sha}"
|
||||
|
||||
|
||||
__version__ = get_version()
|
||||
|
||||
__all__ = ["BASE_VERSION", "__version__", "get_version"]
|
||||
@@ -1,13 +1,14 @@
|
||||
import asyncio
|
||||
import hashlib
|
||||
import secrets
|
||||
import time
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlmodel import select
|
||||
from sqlmodel import col, select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from .core.db import ApiKey, LightningInvoice, get_session
|
||||
from .core.db import ApiKey, LightningInvoice, create_session, get_session
|
||||
from .core.logging import get_logger
|
||||
from .core.settings import settings
|
||||
from .wallet import get_wallet
|
||||
@@ -19,10 +20,27 @@ lightning_router = APIRouter(prefix="/lightning")
|
||||
|
||||
class InvoiceCreateRequest(BaseModel):
|
||||
amount_sats: int = Field(gt=0, le=1_000_000, description="Amount in satoshis")
|
||||
purpose: str = Field(description="create or topup", pattern="^(create|topup)$")
|
||||
api_key: str | None = Field(
|
||||
default=None, description="Required for topup operations"
|
||||
purpose: str = Field(
|
||||
default="create",
|
||||
description="create or topup",
|
||||
pattern="^(create|topup)$",
|
||||
)
|
||||
api_key: str | None = Field(
|
||||
default=None,
|
||||
description="Deprecated: legacy field for topup. Prefer Authorization header.",
|
||||
)
|
||||
balance_limit: int | None = Field(default=None)
|
||||
balance_limit_reset: str | None = Field(default=None)
|
||||
validity_date: int | None = Field(default=None)
|
||||
|
||||
|
||||
def _extract_bearer_api_key(authorization: str | None) -> str | None:
|
||||
if not authorization:
|
||||
return None
|
||||
token = authorization.strip()
|
||||
if token.lower().startswith("bearer "):
|
||||
token = token[7:].strip()
|
||||
return token or None
|
||||
|
||||
|
||||
class InvoiceCreateResponse(BaseModel):
|
||||
@@ -61,18 +79,21 @@ def generate_invoice_id() -> str:
|
||||
@lightning_router.post("/invoice", response_model=InvoiceCreateResponse)
|
||||
async def create_invoice(
|
||||
request: InvoiceCreateRequest,
|
||||
authorization: str | None = Header(default=None),
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> InvoiceCreateResponse:
|
||||
if request.purpose == "topup" and not request.api_key:
|
||||
raise HTTPException(
|
||||
status_code=400, detail="api_key is required for topup operations"
|
||||
)
|
||||
api_key_token = _extract_bearer_api_key(authorization) or request.api_key
|
||||
|
||||
if request.purpose == "topup" and request.api_key:
|
||||
if not request.api_key.startswith("sk-"):
|
||||
if request.purpose == "topup":
|
||||
if not api_key_token:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Authorization bearer api key is required for topup",
|
||||
)
|
||||
if not api_key_token.startswith("sk-"):
|
||||
raise HTTPException(status_code=400, detail="Invalid API key format")
|
||||
|
||||
api_key = await session.get(ApiKey, request.api_key[3:])
|
||||
api_key = await session.get(ApiKey, api_key_token[3:])
|
||||
if not api_key:
|
||||
raise HTTPException(status_code=404, detail="API key not found")
|
||||
|
||||
@@ -92,8 +113,11 @@ async def create_invoice(
|
||||
description=description,
|
||||
payment_hash=payment_hash,
|
||||
status="pending",
|
||||
api_key_hash=request.api_key[3:] if request.api_key else None,
|
||||
api_key_hash=api_key_token[3:] if api_key_token else None,
|
||||
purpose=request.purpose,
|
||||
balance_limit=request.balance_limit,
|
||||
balance_limit_reset=request.balance_limit_reset,
|
||||
validity_date=request.validity_date,
|
||||
expires_at=expires_at,
|
||||
)
|
||||
|
||||
@@ -136,13 +160,13 @@ async def get_invoice_status(
|
||||
if not invoice:
|
||||
raise HTTPException(status_code=404, detail="Invoice not found")
|
||||
|
||||
if invoice.status == "pending":
|
||||
await check_invoice_payment(invoice, session)
|
||||
|
||||
if invoice.status == "pending" and int(time.time()) > invoice.expires_at:
|
||||
invoice.status = "expired"
|
||||
await session.commit()
|
||||
|
||||
if invoice.status == "pending":
|
||||
await check_invoice_payment(invoice, session)
|
||||
|
||||
api_key = None
|
||||
if invoice.status == "paid" and invoice.purpose == "create":
|
||||
if invoice.api_key_hash:
|
||||
@@ -268,3 +292,41 @@ async def topup_api_key_from_invoice(
|
||||
|
||||
api_key.balance += invoice.amount_sats * 1000 # Convert to msats
|
||||
await session.flush()
|
||||
|
||||
|
||||
INVOICE_WATCH_INTERVAL_SECONDS = 5
|
||||
INVOICE_WATCH_BATCH_LIMIT = 100
|
||||
|
||||
|
||||
async def periodic_invoice_watcher() -> None:
|
||||
"""Background task: detect paid Lightning invoices and credit balances.
|
||||
|
||||
Removes the need for clients to poll the status endpoint after paying.
|
||||
"""
|
||||
while True:
|
||||
try:
|
||||
async with create_session() as session:
|
||||
now = int(time.time())
|
||||
result = await session.exec(
|
||||
select(LightningInvoice)
|
||||
.where(
|
||||
LightningInvoice.status == "pending",
|
||||
col(LightningInvoice.expires_at) > now,
|
||||
)
|
||||
.limit(INVOICE_WATCH_BATCH_LIMIT)
|
||||
)
|
||||
pending = result.all()
|
||||
for invoice in pending:
|
||||
try:
|
||||
await check_invoice_payment(invoice, session)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Invoice watcher failed for invoice",
|
||||
extra={"invoice_id": invoice.id, "error": str(e)},
|
||||
)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Invoice watcher loop error: {e}")
|
||||
|
||||
await asyncio.sleep(INVOICE_WATCH_INTERVAL_SECONDS)
|
||||
|
||||
5
routstr/nostr/__init__.py
Normal file
5
routstr/nostr/__init__.py
Normal file
@@ -0,0 +1,5 @@
|
||||
from .analytics import publish_usage_analytics
|
||||
from .discovery import providers_cache_refresher
|
||||
from .listing import announce_provider
|
||||
|
||||
__all__ = ["providers_cache_refresher", "announce_provider", "publish_usage_analytics"]
|
||||
419
routstr/nostr/analytics.py
Normal file
419
routstr/nostr/analytics.py
Normal file
@@ -0,0 +1,419 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Nostr usage analytics publisher.
|
||||
Publishes a single replaceable analytics snapshot for each provider.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from nostr.event import Event
|
||||
from nostr.key import PrivateKey
|
||||
|
||||
from ..core import get_logger
|
||||
from ..core.log_manager import log_manager
|
||||
from ..core.settings import settings
|
||||
from .listing import nsec_to_keypair, publish_to_relay
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
ANALYTICS_KIND = 38422
|
||||
ANALYTICS_SCHEMA = "routstr.analytics.snapshot.v1"
|
||||
DEFAULT_RELAYS = [
|
||||
"wss://relay.nostr.band",
|
||||
"wss://relay.damus.io",
|
||||
"wss://relay.routstr.com",
|
||||
"wss://nos.lol",
|
||||
]
|
||||
PUBLISH_INTERVAL_SECONDS = 15 * 60
|
||||
DISABLED_POLL_SECONDS = 60
|
||||
DASHBOARD_WINDOW_HOURS = 24
|
||||
DASHBOARD_INTERVAL_MINUTES = 60
|
||||
MODEL_LIMIT = 20
|
||||
WINDOW_DEFINITIONS: tuple[tuple[str, int, int], ...] = (
|
||||
("24h", 24, 60),
|
||||
("7d", 7 * 24, 6 * 60),
|
||||
("30d", 30 * 24, 24 * 60),
|
||||
("3m", 90 * 24, 24 * 60),
|
||||
("1y", 365 * 24, 7 * 24 * 60),
|
||||
)
|
||||
|
||||
|
||||
def _event_to_dict(ev: Event) -> dict[str, Any]:
|
||||
return {
|
||||
"id": ev.id,
|
||||
"pubkey": ev.public_key,
|
||||
"created_at": ev.created_at,
|
||||
"kind": int(ev.kind) if not isinstance(ev.kind, int) else ev.kind,
|
||||
"tags": ev.tags,
|
||||
"content": ev.content,
|
||||
"sig": ev.signature,
|
||||
}
|
||||
|
||||
|
||||
def _resolve_provider_id(public_key_hex: str) -> str:
|
||||
explicit_provider_id = (settings.provider_id or "").strip()
|
||||
if explicit_provider_id:
|
||||
return explicit_provider_id
|
||||
return public_key_hex[:12]
|
||||
|
||||
|
||||
def _resolve_endpoint_urls() -> list[str]:
|
||||
urls: list[str] = []
|
||||
http_url = (settings.http_url or "").strip()
|
||||
onion_url = (settings.onion_url or "").strip()
|
||||
|
||||
if http_url and http_url != "http://localhost:8000":
|
||||
urls.append(http_url)
|
||||
|
||||
if onion_url:
|
||||
if onion_url.endswith(".onion") and not (
|
||||
onion_url.startswith("http://") or onion_url.startswith("https://")
|
||||
):
|
||||
onion_url = f"http://{onion_url}"
|
||||
urls.append(onion_url)
|
||||
|
||||
return urls
|
||||
|
||||
|
||||
def _resolve_relays() -> list[str]:
|
||||
configured = [url.strip() for url in settings.relays if url.strip()]
|
||||
return configured if configured else list(DEFAULT_RELAYS)
|
||||
|
||||
|
||||
def _to_int(value: Any) -> int:
|
||||
if isinstance(value, bool):
|
||||
return int(value)
|
||||
if isinstance(value, int):
|
||||
return value
|
||||
if isinstance(value, float):
|
||||
return int(value)
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
return int(float(value))
|
||||
except ValueError:
|
||||
return 0
|
||||
return 0
|
||||
|
||||
|
||||
def _to_float(value: Any) -> float:
|
||||
if isinstance(value, bool):
|
||||
return float(int(value))
|
||||
if isinstance(value, (int, float)):
|
||||
return float(value)
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
return float(value)
|
||||
except ValueError:
|
||||
return 0.0
|
||||
return 0.0
|
||||
|
||||
|
||||
def _aggregate_top_model_usage(
|
||||
model_usage_mix: dict[str, Any],
|
||||
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
|
||||
top_models_raw = model_usage_mix.get("top_models", [])
|
||||
mix_metrics_raw = model_usage_mix.get("metrics", [])
|
||||
|
||||
top_models = [model for model in top_models_raw if isinstance(model, str)]
|
||||
metrics = [row for row in mix_metrics_raw if isinstance(row, dict)]
|
||||
|
||||
model_totals: dict[str, dict[str, float | int]] = {
|
||||
model: {
|
||||
"successful_requests": 0,
|
||||
"revenue_msats": 0.0,
|
||||
"total_tokens": 0,
|
||||
}
|
||||
for model in top_models
|
||||
}
|
||||
others = {
|
||||
"successful_requests": 0,
|
||||
"revenue_msats": 0.0,
|
||||
"total_tokens": 0,
|
||||
}
|
||||
|
||||
for metric in metrics:
|
||||
model_counts = metric.get("model_counts", {})
|
||||
model_revenue = metric.get("model_revenue_msats", {})
|
||||
model_tokens = metric.get("model_tokens", {})
|
||||
|
||||
if isinstance(model_counts, dict):
|
||||
for model, count in model_counts.items():
|
||||
if model in model_totals:
|
||||
model_totals[model]["successful_requests"] += _to_int(count)
|
||||
|
||||
if isinstance(model_revenue, dict):
|
||||
for model, amount in model_revenue.items():
|
||||
if model in model_totals:
|
||||
model_totals[model]["revenue_msats"] += _to_float(amount)
|
||||
|
||||
if isinstance(model_tokens, dict):
|
||||
for model, token_count in model_tokens.items():
|
||||
if model in model_totals:
|
||||
model_totals[model]["total_tokens"] += _to_int(token_count)
|
||||
|
||||
others["successful_requests"] += _to_int(metric.get("others", 0))
|
||||
others["revenue_msats"] += _to_float(metric.get("others_revenue_msats", 0.0))
|
||||
others["total_tokens"] += _to_int(metric.get("others_tokens", 0))
|
||||
|
||||
model_rows = [
|
||||
{
|
||||
"model": model,
|
||||
"successful_requests": int(values["successful_requests"]),
|
||||
"revenue_msats": float(values["revenue_msats"]),
|
||||
"total_tokens": int(values["total_tokens"]),
|
||||
}
|
||||
for model, values in model_totals.items()
|
||||
]
|
||||
model_rows.sort(
|
||||
key=lambda row: _to_int(row.get("successful_requests", 0)),
|
||||
reverse=True,
|
||||
)
|
||||
|
||||
return model_rows, others
|
||||
|
||||
|
||||
def _build_summary_payload(summary: dict[str, Any]) -> dict[str, Any]:
|
||||
return {
|
||||
"total_requests": _to_int(summary.get("total_requests", 0)),
|
||||
"successful_chat_completions": _to_int(
|
||||
summary.get("successful_chat_completions", 0)
|
||||
),
|
||||
"failed_requests": _to_int(summary.get("failed_requests", 0)),
|
||||
"success_rate": _to_float(summary.get("success_rate", 0.0)),
|
||||
"unique_models_count": _to_int(summary.get("unique_models_count", 0)),
|
||||
"input_tokens": _to_int(summary.get("input_tokens", 0)),
|
||||
"output_tokens": _to_int(summary.get("output_tokens", 0)),
|
||||
"total_tokens": _to_int(summary.get("total_tokens", 0)),
|
||||
"revenue_msats": _to_float(summary.get("revenue_msats", 0.0)),
|
||||
"refunds_msats": _to_float(summary.get("refunds_msats", 0.0)),
|
||||
"net_revenue_msats": _to_float(summary.get("net_revenue_msats", 0.0)),
|
||||
"revenue_sats": _to_float(summary.get("revenue_sats", 0.0)),
|
||||
"refunds_sats": _to_float(summary.get("refunds_sats", 0.0)),
|
||||
"net_revenue_sats": _to_float(summary.get("net_revenue_sats", 0.0)),
|
||||
}
|
||||
|
||||
|
||||
def _build_window_payload(
|
||||
*,
|
||||
hours: int,
|
||||
interval_minutes: int,
|
||||
model_limit: int,
|
||||
) -> dict[str, Any]:
|
||||
dashboard = log_manager.get_usage_dashboard(
|
||||
interval=interval_minutes,
|
||||
hours=hours,
|
||||
error_limit=1,
|
||||
model_limit=model_limit,
|
||||
)
|
||||
|
||||
summary = dashboard.get("summary", {})
|
||||
model_usage_mix = dashboard.get("model_usage_mix", {})
|
||||
|
||||
summary_payload = _build_summary_payload(summary if isinstance(summary, dict) else {})
|
||||
usage_mix_payload = model_usage_mix if isinstance(model_usage_mix, dict) else {}
|
||||
top_model_usage, others_usage = _aggregate_top_model_usage(usage_mix_payload)
|
||||
|
||||
return {
|
||||
"window_hours": hours,
|
||||
"interval_minutes": interval_minutes,
|
||||
"summary": summary_payload,
|
||||
"model_usage_mix": usage_mix_payload,
|
||||
"top_model_usage": top_model_usage,
|
||||
"others_usage": others_usage,
|
||||
}
|
||||
|
||||
|
||||
def build_stats_snapshot_payload(
|
||||
provider_id: str,
|
||||
*,
|
||||
public_key_hex: str,
|
||||
generated_at: int,
|
||||
window_hours: int = DASHBOARD_WINDOW_HOURS,
|
||||
interval_minutes: int = DASHBOARD_INTERVAL_MINUTES,
|
||||
model_limit: int = MODEL_LIMIT,
|
||||
) -> dict[str, Any]:
|
||||
_ = (window_hours, interval_minutes)
|
||||
windows: dict[str, dict[str, Any]] = {}
|
||||
for key, hours, window_interval_minutes in WINDOW_DEFINITIONS:
|
||||
windows[key] = _build_window_payload(
|
||||
hours=hours,
|
||||
interval_minutes=window_interval_minutes,
|
||||
model_limit=model_limit,
|
||||
)
|
||||
|
||||
primary_window = windows.get("24h", {})
|
||||
summary_payload = (
|
||||
primary_window.get("summary", {})
|
||||
if isinstance(primary_window.get("summary", {}), dict)
|
||||
else {}
|
||||
)
|
||||
usage_mix_payload = (
|
||||
primary_window.get("model_usage_mix", {})
|
||||
if isinstance(primary_window.get("model_usage_mix", {}), dict)
|
||||
else {}
|
||||
)
|
||||
top_model_usage = (
|
||||
primary_window.get("top_model_usage", [])
|
||||
if isinstance(primary_window.get("top_model_usage", []), list)
|
||||
else []
|
||||
)
|
||||
others_usage = (
|
||||
primary_window.get("others_usage", {})
|
||||
if isinstance(primary_window.get("others_usage", {}), dict)
|
||||
else {}
|
||||
)
|
||||
|
||||
return {
|
||||
"schema": ANALYTICS_SCHEMA,
|
||||
"generated_at": generated_at,
|
||||
"provider_id": provider_id,
|
||||
"pubkey": public_key_hex,
|
||||
"npub": settings.npub or "",
|
||||
"endpoint_urls": _resolve_endpoint_urls(),
|
||||
"window_hours": DASHBOARD_WINDOW_HOURS,
|
||||
"interval_minutes": DASHBOARD_INTERVAL_MINUTES,
|
||||
"summary": summary_payload,
|
||||
"model_usage_mix": usage_mix_payload,
|
||||
"top_model_usage": top_model_usage,
|
||||
"others_usage": others_usage,
|
||||
"windows": windows,
|
||||
}
|
||||
|
||||
|
||||
def create_stats_snapshot_event(
|
||||
private_key_hex: str,
|
||||
provider_id: str,
|
||||
payload_json: str,
|
||||
*,
|
||||
d_tag: str,
|
||||
) -> dict[str, Any]:
|
||||
private_key = PrivateKey(bytes.fromhex(private_key_hex))
|
||||
tags = [
|
||||
["d", d_tag],
|
||||
["provider", provider_id],
|
||||
["schema", ANALYTICS_SCHEMA],
|
||||
]
|
||||
|
||||
event = Event(
|
||||
public_key=private_key.public_key.hex(),
|
||||
content=payload_json,
|
||||
kind=ANALYTICS_KIND,
|
||||
tags=tags,
|
||||
)
|
||||
private_key.sign_event(event)
|
||||
return _event_to_dict(event)
|
||||
|
||||
|
||||
def _fingerprint_payload(payload: dict[str, Any]) -> str:
|
||||
normalized = dict(payload)
|
||||
# Ignore generated timestamp for semantic dedupe.
|
||||
normalized.pop("generated_at", None)
|
||||
payload_json = json.dumps(normalized, separators=(",", ":"), sort_keys=True)
|
||||
return hashlib.sha256(payload_json.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
async def publish_usage_analytics() -> None:
|
||||
last_payload_hash: str | None = None
|
||||
|
||||
parsed_nsec: str | None = None
|
||||
private_key_hex: str | None = None
|
||||
public_key_hex: str | None = None
|
||||
provider_id: str | None = None
|
||||
warned_missing_nsec = False
|
||||
|
||||
logger.info("Usage analytics sharing task started")
|
||||
|
||||
while True:
|
||||
try:
|
||||
if not settings.enable_analytics_sharing:
|
||||
await asyncio.sleep(DISABLED_POLL_SECONDS)
|
||||
continue
|
||||
|
||||
nsec = (settings.nsec or "").strip()
|
||||
if not nsec:
|
||||
if not warned_missing_nsec:
|
||||
logger.info("NSEC is not configured; skipping analytics sharing to Nostr")
|
||||
warned_missing_nsec = True
|
||||
await asyncio.sleep(DISABLED_POLL_SECONDS)
|
||||
continue
|
||||
|
||||
warned_missing_nsec = False
|
||||
if nsec != parsed_nsec or private_key_hex is None or public_key_hex is None:
|
||||
keypair = nsec_to_keypair(nsec)
|
||||
if not keypair:
|
||||
logger.error("Invalid NSEC; analytics sharing is paused")
|
||||
await asyncio.sleep(DISABLED_POLL_SECONDS)
|
||||
continue
|
||||
private_key_hex, public_key_hex = keypair
|
||||
parsed_nsec = nsec
|
||||
provider_id = _resolve_provider_id(public_key_hex)
|
||||
last_payload_hash = None
|
||||
|
||||
if private_key_hex is None or public_key_hex is None:
|
||||
await asyncio.sleep(DISABLED_POLL_SECONDS)
|
||||
continue
|
||||
|
||||
relay_urls = _resolve_relays()
|
||||
if not relay_urls:
|
||||
logger.warning("No Nostr relays configured; analytics sharing skipped")
|
||||
await asyncio.sleep(DISABLED_POLL_SECONDS)
|
||||
continue
|
||||
|
||||
resolved_provider_id = provider_id or _resolve_provider_id(public_key_hex)
|
||||
now_ts = int(time.time())
|
||||
payload = build_stats_snapshot_payload(
|
||||
resolved_provider_id,
|
||||
public_key_hex=public_key_hex,
|
||||
generated_at=now_ts,
|
||||
)
|
||||
|
||||
payload_hash = _fingerprint_payload(payload)
|
||||
if last_payload_hash == payload_hash:
|
||||
await asyncio.sleep(PUBLISH_INTERVAL_SECONDS)
|
||||
continue
|
||||
|
||||
payload_json = json.dumps(payload, separators=(",", ":"), sort_keys=True)
|
||||
d_tag = f"{resolved_provider_id}:stats"
|
||||
event = create_stats_snapshot_event(
|
||||
private_key_hex,
|
||||
resolved_provider_id,
|
||||
payload_json,
|
||||
d_tag=d_tag,
|
||||
)
|
||||
|
||||
success_count = 0
|
||||
for relay_url in relay_urls:
|
||||
if await publish_to_relay(relay_url, event):
|
||||
success_count += 1
|
||||
|
||||
if success_count > 0:
|
||||
last_payload_hash = payload_hash
|
||||
|
||||
logger.info(
|
||||
"Published analytics snapshot (success=%s/%s provider=%s)",
|
||||
success_count,
|
||||
len(relay_urls),
|
||||
resolved_provider_id,
|
||||
extra={
|
||||
"relay_success_count": success_count,
|
||||
"relay_total": len(relay_urls),
|
||||
"provider_id": resolved_provider_id,
|
||||
},
|
||||
)
|
||||
await asyncio.sleep(PUBLISH_INTERVAL_SECONDS)
|
||||
|
||||
except asyncio.CancelledError:
|
||||
logger.info("Usage analytics sharing task cancelled")
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Usage analytics sharing error",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
)
|
||||
await asyncio.sleep(DISABLED_POLL_SECONDS)
|
||||
@@ -8,8 +8,8 @@ import httpx
|
||||
import websockets
|
||||
from fastapi import APIRouter, HTTPException
|
||||
|
||||
from .core.logging import get_logger
|
||||
from .core.settings import settings
|
||||
from ..core.logging import get_logger
|
||||
from ..core.settings import settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
@@ -72,8 +72,6 @@ async def query_nostr_relay_for_providers(
|
||||
elif data[0] == "NOTICE":
|
||||
try:
|
||||
msg = str(data[1])
|
||||
if len(msg) > 200:
|
||||
msg = msg[:200] + "..."
|
||||
logger.debug(f"Relay notice: {msg}")
|
||||
except Exception:
|
||||
logger.debug("Relay notice received")
|
||||
@@ -1,6 +1,6 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
NIP-91: Routstr Provider Discoverability Implementation
|
||||
Listing: Routstr Provider Discoverability Implementation
|
||||
Automatically announces this Routstr proxy instance to Nostr relays.
|
||||
"""
|
||||
|
||||
@@ -18,15 +18,15 @@ from nostr.key import PrivateKey
|
||||
from nostr.message_type import ClientMessageType
|
||||
from nostr.relay_manager import RelayManager
|
||||
|
||||
from .core import get_logger
|
||||
from .core.settings import settings
|
||||
from ..core import get_logger
|
||||
from ..core.settings import settings
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
def get_app_version() -> str | None:
|
||||
try:
|
||||
from .core.main import __version__ as imported_version
|
||||
from ..core.version import __version__ as imported_version
|
||||
|
||||
return imported_version
|
||||
except Exception:
|
||||
@@ -71,7 +71,7 @@ def nsec_to_keypair(nsec: str) -> tuple[str, str] | None:
|
||||
return None
|
||||
|
||||
|
||||
def create_nip91_event(
|
||||
def create_listing_event(
|
||||
private_key_hex: str,
|
||||
provider_id: str,
|
||||
endpoint_urls: list[str],
|
||||
@@ -80,7 +80,7 @@ def create_nip91_event(
|
||||
metadata: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Create a NIP-91 compliant provider announcement event (kind:38421).
|
||||
Create a listing provider announcement event (kind:38421).
|
||||
|
||||
Args:
|
||||
private_key_hex: 32-byte hex private key for signing
|
||||
@@ -164,14 +164,14 @@ def events_semantically_equal(a: dict[str, Any], b: dict[str, Any]) -> bool:
|
||||
return True
|
||||
|
||||
|
||||
async def query_nip91_events(
|
||||
async def query_listing_events(
|
||||
relay_url: str,
|
||||
pubkey: str,
|
||||
provider_id: str | None = None,
|
||||
timeout: int = 30,
|
||||
) -> tuple[list[dict[str, Any]], bool]:
|
||||
"""
|
||||
Query a Nostr relay for NIP-91 provider announcements (kind:38421) via nostr library.
|
||||
Query a Nostr relay for listing provider announcements (kind:38421) via nostr library.
|
||||
|
||||
Returns a tuple of (events, ok) where ok indicates whether the relay interaction
|
||||
succeeded without transport-level errors.
|
||||
@@ -188,7 +188,7 @@ async def query_nip91_events(
|
||||
|
||||
flt = Filter(kinds=[38421], authors=[pubkey], limit=10)
|
||||
filters = Filters([flt])
|
||||
sub_id = f"nip91_{int(time.time())}"
|
||||
sub_id = f"routstr_listing_{int(time.time())}"
|
||||
rm.add_subscription(sub_id, filters)
|
||||
req: list[Any] = [ClientMessageType.REQUEST, sub_id]
|
||||
req.extend(filters.to_json_array())
|
||||
@@ -294,7 +294,7 @@ async def _determine_provider_id(public_key_hex: str, relay_urls: list[str]) ->
|
||||
|
||||
async def query_single_relay(relay_url: str) -> list[dict[str, Any]]:
|
||||
try:
|
||||
events, _ok = await query_nip91_events(relay_url, public_key_hex, None)
|
||||
events, _ok = await query_listing_events(relay_url, public_key_hex, None)
|
||||
return events
|
||||
except Exception:
|
||||
return []
|
||||
@@ -330,7 +330,7 @@ async def publish_to_relay(
|
||||
timeout: int = 30,
|
||||
) -> bool:
|
||||
"""
|
||||
Publish a NIP-91 event to a nostr relay via nostr library.
|
||||
Publish a listing event to a nostr relay via nostr library.
|
||||
"""
|
||||
|
||||
def _sync_publish() -> bool:
|
||||
@@ -341,7 +341,7 @@ async def publish_to_relay(
|
||||
time.sleep(1.0)
|
||||
# Publish the event as-is via publish_message to preserve signature
|
||||
rm.publish_message(json.dumps(["EVENT", event]))
|
||||
logger.debug(f"Sent NIP-91 event {event.get('id', '')} to {relay_url}")
|
||||
logger.debug(f"Sent listing event {event.get('id', '')} to {relay_url}")
|
||||
time.sleep(1.0)
|
||||
return True
|
||||
except Exception as e:
|
||||
@@ -364,13 +364,13 @@ async def announce_provider() -> None:
|
||||
# Check for NSEC in environment (use NSEC only)
|
||||
nsec = settings.nsec
|
||||
if not nsec:
|
||||
logger.info("Nostr private key not found (NSEC), skipping NIP-91 announcement")
|
||||
logger.info("Nostr private key not found (NSEC), skipping listing announcement")
|
||||
return
|
||||
|
||||
# Convert NSEC to keypair
|
||||
keypair = nsec_to_keypair(nsec)
|
||||
if not keypair:
|
||||
logger.error("Failed to parse NSEC, skipping NIP-91 announcement")
|
||||
logger.error("Failed to parse NSEC, skipping listing announcement")
|
||||
return
|
||||
|
||||
private_key_hex, public_key_hex = keypair
|
||||
@@ -409,7 +409,7 @@ async def announce_provider() -> None:
|
||||
|
||||
if not endpoint_urls:
|
||||
logger.warning(
|
||||
"No valid endpoints configured (HTTP_URL/ONION_URL). Skipping NIP-91 publish."
|
||||
"No valid endpoints configured (HTTP_URL/ONION_URL). Skipping listing publish."
|
||||
)
|
||||
return
|
||||
|
||||
@@ -434,7 +434,7 @@ async def announce_provider() -> None:
|
||||
|
||||
# Create the candidate event that we would publish
|
||||
version_str = get_app_version()
|
||||
candidate_event = create_nip91_event(
|
||||
candidate_event = create_listing_event(
|
||||
private_key_hex=private_key_hex,
|
||||
provider_id=provider_id,
|
||||
endpoint_urls=endpoint_urls,
|
||||
@@ -474,7 +474,7 @@ async def announce_provider() -> None:
|
||||
if _should_skip(relay_url):
|
||||
logger.debug(f"Skipping {relay_url} due to backoff")
|
||||
continue
|
||||
events, ok = await query_nip91_events(relay_url, public_key_hex, provider_id)
|
||||
events, ok = await query_listing_events(relay_url, public_key_hex, provider_id)
|
||||
if ok:
|
||||
_register_success(relay_url)
|
||||
existing_events.extend(events)
|
||||
@@ -489,7 +489,7 @@ async def announce_provider() -> None:
|
||||
|
||||
if not all_match:
|
||||
logger.debug(
|
||||
"No matching NIP-91 announcement found or differences detected; publishing update"
|
||||
"No matching listing announcement found or differences detected; publishing update"
|
||||
)
|
||||
success_count = 0
|
||||
for relay_url in relay_urls:
|
||||
@@ -502,11 +502,11 @@ async def announce_provider() -> None:
|
||||
else:
|
||||
_register_failure(relay_url)
|
||||
logger.info(
|
||||
f"Published NIP-91 announcement to {success_count}/{len(relay_urls)} relays"
|
||||
f"Published listing announcement to {success_count}/{len(relay_urls)} relays"
|
||||
)
|
||||
else:
|
||||
logger.debug(
|
||||
"Matching NIP-91 announcement already present; skipping publish on startup"
|
||||
"Matching listing announcement already present; skipping publish on startup"
|
||||
)
|
||||
|
||||
# Re-announce periodically (every 24 hours)
|
||||
@@ -518,7 +518,7 @@ async def announce_provider() -> None:
|
||||
|
||||
# Build fresh candidate event for comparison
|
||||
version_str = get_app_version()
|
||||
candidate_event = create_nip91_event(
|
||||
candidate_event = create_listing_event(
|
||||
private_key_hex=private_key_hex,
|
||||
provider_id=provider_id,
|
||||
endpoint_urls=endpoint_urls,
|
||||
@@ -533,7 +533,7 @@ async def announce_provider() -> None:
|
||||
if _should_skip(relay_url):
|
||||
logger.debug(f"Skipping {relay_url} due to backoff")
|
||||
continue
|
||||
events, ok = await query_nip91_events(
|
||||
events, ok = await query_listing_events(
|
||||
relay_url, public_key_hex, provider_id
|
||||
)
|
||||
if ok:
|
||||
@@ -549,7 +549,7 @@ async def announce_provider() -> None:
|
||||
|
||||
if all_match:
|
||||
logger.debug(
|
||||
"Matching NIP-91 announcement already present; skipping periodic re-announce"
|
||||
"Matching listing announcement already present; skipping periodic re-announce"
|
||||
)
|
||||
continue
|
||||
|
||||
@@ -567,8 +567,8 @@ async def announce_provider() -> None:
|
||||
_register_failure(relay_url)
|
||||
|
||||
except asyncio.CancelledError:
|
||||
logger.info("NIP-91 announcement task cancelled")
|
||||
logger.info("Listing announcement task cancelled")
|
||||
break
|
||||
except Exception as e:
|
||||
logger.debug(f"Error in NIP-91 announcement loop: {type(e).__name__}")
|
||||
logger.debug(f"Error in listing announcement loop: {type(e).__name__}")
|
||||
# Continue running despite errors
|
||||
@@ -15,6 +15,13 @@ class CostData(BaseModel):
|
||||
input_msats: int
|
||||
output_msats: int
|
||||
total_msats: int
|
||||
total_usd: float = 0.0
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
cache_read_input_tokens: int = 0
|
||||
cache_creation_input_tokens: int = 0
|
||||
cache_read_msats: int = 0
|
||||
cache_creation_msats: int = 0
|
||||
|
||||
|
||||
class MaxCostData(CostData):
|
||||
@@ -26,11 +33,10 @@ class CostDataError(BaseModel):
|
||||
code: str
|
||||
|
||||
|
||||
async def calculate_cost( # todo: can be sync
|
||||
async def calculate_cost(
|
||||
response_data: dict, max_cost: int, session: AsyncSession
|
||||
) -> CostData | MaxCostData | CostDataError:
|
||||
"""
|
||||
Calculate the cost of an API request based on token usage.
|
||||
"""Calculate the cost of an API request based on token usage.
|
||||
|
||||
Args:
|
||||
response_data: Response data containing usage information
|
||||
@@ -48,12 +54,20 @@ async def calculate_cost( # todo: can be sync
|
||||
},
|
||||
)
|
||||
|
||||
# Check for usage data
|
||||
if "usage" not in response_data or response_data["usage"] is None:
|
||||
logger.warning(
|
||||
"No usage data in response, using base cost only",
|
||||
"No usage data in response — billing at MaxCostData with zero "
|
||||
"tokens. Dashboard will show this request as `(0+0)`. Most "
|
||||
"common cause: upstream stream did not include a final usage "
|
||||
"chunk (OpenAI-compat backends require "
|
||||
"`stream_options.include_usage=true`).",
|
||||
extra={
|
||||
"max_cost_msats": max_cost,
|
||||
"model": response_data.get("model", "unknown"),
|
||||
"response_keys": sorted(response_data.keys())
|
||||
if isinstance(response_data, dict)
|
||||
else None,
|
||||
},
|
||||
)
|
||||
return MaxCostData(
|
||||
@@ -61,46 +75,61 @@ async def calculate_cost( # todo: can be sync
|
||||
input_msats=0,
|
||||
output_msats=0,
|
||||
total_msats=0,
|
||||
total_usd=0.0,
|
||||
input_tokens=0,
|
||||
output_tokens=0,
|
||||
cache_read_input_tokens=0,
|
||||
cache_creation_input_tokens=0,
|
||||
cache_read_msats=0,
|
||||
cache_creation_msats=0,
|
||||
)
|
||||
|
||||
usage_data = response_data["usage"]
|
||||
|
||||
usd_cost = 0.0
|
||||
# Extract token counts
|
||||
input_tokens = _extract_token_pair(usage_data, "prompt_tokens", "input_tokens")
|
||||
output_tokens = _extract_token_pair(usage_data, "completion_tokens", "output_tokens")
|
||||
|
||||
# Prioritize cost_details.upstream_inference_cost
|
||||
if "cost_details" in usage_data:
|
||||
usd_cost = float(
|
||||
usage_data["cost_details"].get("upstream_inference_cost", 0) or 0
|
||||
)
|
||||
|
||||
# Fallback to cost field if upstream_inference_cost is 0
|
||||
if usd_cost == 0 and "cost" in usage_data:
|
||||
try:
|
||||
usd_cost = float(usage_data.get("cost", 0) or 0)
|
||||
except Exception:
|
||||
pass
|
||||
# Extract cache tokens (handles OpenAI vs Anthropic formats)
|
||||
cache_read_tokens, cache_creation_tokens, input_tokens = _extract_cache_tokens(
|
||||
usage_data, input_tokens
|
||||
)
|
||||
|
||||
# Try USD cost first
|
||||
usd_cost = _resolve_usd_cost(usage_data, response_data)
|
||||
if usd_cost > 0:
|
||||
try:
|
||||
sats_per_usd = 1.0 / sats_usd_price()
|
||||
cost_in_sats = usd_cost * sats_per_usd
|
||||
cost_in_msats = math.ceil(cost_in_sats * 1000)
|
||||
|
||||
logger.info(
|
||||
"Using cost from usage data/details",
|
||||
if input_tokens == 0 and output_tokens == 0:
|
||||
logger.warning(
|
||||
"Upstream reported a USD cost but no token counts — "
|
||||
"billing the USD-derived cost while the dashboard will "
|
||||
"show this request as `(0+0)` tokens. Check that the "
|
||||
"upstream actually emits `usage.input_tokens` and "
|
||||
"`usage.output_tokens` (OpenAI-compat streams require "
|
||||
"`stream_options.include_usage=true`).",
|
||||
extra={
|
||||
"usd_cost": usd_cost,
|
||||
"cost_in_sats": cost_in_sats,
|
||||
"cost_in_msats": cost_in_msats,
|
||||
"model": response_data.get("model", "unknown"),
|
||||
"usd_cost": usd_cost,
|
||||
"usage_keys": sorted(usage_data.keys())
|
||||
if isinstance(usage_data, dict)
|
||||
else None,
|
||||
},
|
||||
)
|
||||
|
||||
return CostData(
|
||||
base_msats=-1,
|
||||
input_msats=-1, # Cost field doesn't break down by token type
|
||||
output_msats=-1,
|
||||
total_msats=cost_in_msats,
|
||||
try:
|
||||
input_usd = _coerce_usd(
|
||||
usage_data.get("cost_details", {}).get("input_cost", 0)
|
||||
)
|
||||
output_usd = _coerce_usd(
|
||||
usage_data.get("cost_details", {}).get("output_cost", 0)
|
||||
)
|
||||
return _calculate_from_usd_cost(
|
||||
usd_cost,
|
||||
input_usd,
|
||||
output_usd,
|
||||
input_tokens,
|
||||
cache_read_tokens,
|
||||
cache_creation_tokens,
|
||||
output_tokens,
|
||||
response_data,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
@@ -111,69 +140,35 @@ async def calculate_cost( # todo: can be sync
|
||||
"model": response_data.get("model", "unknown"),
|
||||
},
|
||||
)
|
||||
# Fall through to token-based calculation
|
||||
|
||||
MSATS_PER_1K_INPUT_TOKENS: float = (
|
||||
float(settings.fixed_per_1k_input_tokens) * 1000.0
|
||||
)
|
||||
MSATS_PER_1K_OUTPUT_TOKENS: float = (
|
||||
float(settings.fixed_per_1k_output_tokens) * 1000.0
|
||||
)
|
||||
# Fall back to token-based pricing
|
||||
try:
|
||||
pricing_rates = _get_pricing_rates(response_data)
|
||||
except ValueError as e:
|
||||
return CostDataError(message=str(e), code="pricing_error")
|
||||
|
||||
if not settings.fixed_pricing:
|
||||
response_model = response_data.get("model", "")
|
||||
logger.debug(
|
||||
"Using model-based pricing",
|
||||
extra={"model": response_model},
|
||||
)
|
||||
if pricing_rates is None:
|
||||
input_rate = float(settings.fixed_per_1k_input_tokens) * 1000.0
|
||||
output_rate = float(settings.fixed_per_1k_output_tokens) * 1000.0
|
||||
cache_read_rate = input_rate
|
||||
cache_creation_rate = input_rate
|
||||
else:
|
||||
input_rate, output_rate, cache_read_rate, cache_creation_rate = pricing_rates
|
||||
|
||||
from ..proxy import get_model_instance
|
||||
|
||||
model_obj = get_model_instance(response_model)
|
||||
|
||||
if not model_obj:
|
||||
logger.error(
|
||||
"Invalid model in response",
|
||||
extra={"response_model": response_model},
|
||||
)
|
||||
return CostDataError(
|
||||
message=f"Invalid model in response: {response_model}",
|
||||
code="model_not_found",
|
||||
)
|
||||
|
||||
if not model_obj.sats_pricing:
|
||||
logger.error(
|
||||
"Model pricing not defined",
|
||||
extra={"model": response_model, "model_id": response_model},
|
||||
)
|
||||
return CostDataError(
|
||||
message="Model pricing not defined", code="pricing_not_found"
|
||||
)
|
||||
|
||||
try:
|
||||
mspp = float(model_obj.sats_pricing.prompt)
|
||||
mspc = float(model_obj.sats_pricing.completion)
|
||||
except Exception:
|
||||
return CostDataError(message="Invalid pricing data", code="pricing_invalid")
|
||||
|
||||
MSATS_PER_1K_INPUT_TOKENS = mspp * 1_000_000.0
|
||||
MSATS_PER_1K_OUTPUT_TOKENS = mspc * 1_000_000.0
|
||||
|
||||
logger.info(
|
||||
"Applied model-specific pricing",
|
||||
extra={
|
||||
"model": response_model,
|
||||
"input_price_msats_per_1k": MSATS_PER_1K_INPUT_TOKENS,
|
||||
"output_price_msats_per_1k": MSATS_PER_1K_OUTPUT_TOKENS,
|
||||
},
|
||||
)
|
||||
|
||||
if not (MSATS_PER_1K_OUTPUT_TOKENS and MSATS_PER_1K_INPUT_TOKENS):
|
||||
if not (input_rate and output_rate):
|
||||
logger.warning(
|
||||
"No token pricing configured, using base cost",
|
||||
"No token pricing configured — billing at flat MaxCostData. "
|
||||
"Token counts %s in the upstream response but cannot be "
|
||||
"priced; the request will appear in dashboards with the "
|
||||
"raw counts and a fixed max-cost charge.",
|
||||
"are present"
|
||||
if (input_tokens > 0 or output_tokens > 0)
|
||||
else "are zero",
|
||||
extra={
|
||||
"base_cost_msats": max_cost,
|
||||
"model": response_data.get("model", "unknown"),
|
||||
"input_tokens": input_tokens,
|
||||
"output_tokens": output_tokens,
|
||||
},
|
||||
)
|
||||
return MaxCostData(
|
||||
@@ -181,51 +176,294 @@ async def calculate_cost( # todo: can be sync
|
||||
input_msats=0,
|
||||
output_msats=0,
|
||||
total_msats=max_cost,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cache_read_input_tokens=cache_read_tokens,
|
||||
cache_creation_input_tokens=cache_creation_tokens,
|
||||
cache_read_msats=0,
|
||||
cache_creation_msats=0,
|
||||
)
|
||||
|
||||
input_tokens = usage_data.get("prompt_tokens", 0)
|
||||
output_tokens = usage_data.get("completion_tokens", 0)
|
||||
|
||||
# added for response api
|
||||
input_tokens = (
|
||||
input_tokens if input_tokens != 0 else usage_data.get("input_tokens", 0)
|
||||
)
|
||||
output_tokens = (
|
||||
output_tokens if output_tokens != 0 else usage_data.get("output_tokens", 0)
|
||||
return _calculate_from_tokens(
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
cache_read_tokens,
|
||||
cache_creation_tokens,
|
||||
input_rate,
|
||||
output_rate,
|
||||
cache_read_rate,
|
||||
cache_creation_rate,
|
||||
response_data,
|
||||
)
|
||||
|
||||
# added for response api
|
||||
input_tokens = (
|
||||
input_tokens
|
||||
if input_tokens != 0
|
||||
else response_data.get("usage", {}).get("input_tokens", 0)
|
||||
)
|
||||
output_tokens = (
|
||||
output_tokens
|
||||
if output_tokens != 0
|
||||
else response_data.get("usage", {}).get("output_tokens", 0)
|
||||
|
||||
# ============================================================================
|
||||
# Helper Functions (ordered by call sequence in calculate_cost)
|
||||
# ============================================================================
|
||||
|
||||
|
||||
def parse_token_count(value: object) -> int:
|
||||
"""Parse a token count from various formats (int, float, str, bool)."""
|
||||
if isinstance(value, bool):
|
||||
return 0
|
||||
if isinstance(value, int):
|
||||
return max(0, value)
|
||||
if isinstance(value, float):
|
||||
return max(0, int(value))
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
return max(0, int(float(value)))
|
||||
except ValueError:
|
||||
return 0
|
||||
return 0
|
||||
|
||||
|
||||
def _coerce_usd(value: object) -> float:
|
||||
"""Coerce a value to USD float, handling various formats safely."""
|
||||
if value is None or isinstance(value, bool):
|
||||
return 0.0
|
||||
if not isinstance(value, (int, float, str)):
|
||||
return 0.0
|
||||
try:
|
||||
return max(0.0, float(value))
|
||||
except (TypeError, ValueError):
|
||||
return 0.0
|
||||
|
||||
|
||||
def _extract_token_pair(
|
||||
usage_data: dict, standard_field: str, alt_field: str
|
||||
) -> int:
|
||||
"""Extract token count trying two field names in order."""
|
||||
value = parse_token_count(usage_data.get(standard_field, 0))
|
||||
if value > 0:
|
||||
return value
|
||||
return parse_token_count(usage_data.get(alt_field, 0))
|
||||
|
||||
|
||||
def _extract_cache_tokens(usage_data: dict, input_tokens: int) -> tuple[int, int, int]:
|
||||
"""Extract cache tokens, handling OpenAI vs Anthropic formats.
|
||||
|
||||
Returns: (cache_read_tokens, cache_creation_tokens, adjusted_input_tokens)
|
||||
"""
|
||||
cache_read = parse_token_count(usage_data.get("cache_read_input_tokens", 0))
|
||||
cache_creation = parse_token_count(
|
||||
usage_data.get("cache_creation_input_tokens", 0)
|
||||
)
|
||||
|
||||
input_msats = round(input_tokens / 1000 * MSATS_PER_1K_INPUT_TOKENS, 3)
|
||||
# OpenAI: cache is included in input_tokens, subtract it
|
||||
prompt_details = usage_data.get("prompt_tokens_details")
|
||||
if isinstance(prompt_details, dict) and not cache_read:
|
||||
openai_cached = parse_token_count(prompt_details.get("cached_tokens", 0))
|
||||
if openai_cached:
|
||||
cache_read = openai_cached
|
||||
input_tokens = max(0, input_tokens - cache_read)
|
||||
|
||||
output_msats = round(output_tokens / 1000 * MSATS_PER_1K_OUTPUT_TOKENS, 3)
|
||||
token_based_cost = math.ceil(input_msats + output_msats)
|
||||
return cache_read, cache_creation, input_tokens
|
||||
|
||||
|
||||
def _resolve_usd_cost(usage_data: dict, response_data: dict) -> float:
|
||||
"""Resolve USD cost with clear priority order.
|
||||
|
||||
Priority: cost_details.total_cost → total_cost → cost (in both usage and response).
|
||||
"""
|
||||
cost_details = usage_data.get("cost_details")
|
||||
if isinstance(cost_details, dict):
|
||||
cost = _coerce_usd(cost_details.get("total_cost"))
|
||||
if cost > 0:
|
||||
return cost
|
||||
|
||||
for source in [usage_data, response_data]:
|
||||
if not isinstance(source, dict):
|
||||
continue
|
||||
for field in ("total_cost", "cost"):
|
||||
cost = _coerce_usd(source.get(field))
|
||||
if cost > 0:
|
||||
return cost
|
||||
|
||||
return 0.0
|
||||
|
||||
|
||||
def _get_pricing_rates(
|
||||
response_data: dict,
|
||||
) -> tuple[float, float, float, float] | None:
|
||||
"""Get model-based pricing rates or None if using fixed pricing.
|
||||
|
||||
Returns: (input_rate, output_rate, cache_read_rate, cache_write_rate)
|
||||
"""
|
||||
if settings.fixed_pricing:
|
||||
return None
|
||||
|
||||
from ..proxy import get_model_instance
|
||||
|
||||
response_model = response_data.get("model", "")
|
||||
model_obj = get_model_instance(response_model)
|
||||
|
||||
if not model_obj:
|
||||
logger.error("Invalid model in response", extra={"response_model": response_model})
|
||||
raise ValueError(f"Invalid model: {response_model}")
|
||||
|
||||
if not model_obj.sats_pricing:
|
||||
logger.error(
|
||||
"Model pricing not defined",
|
||||
extra={"model": response_model, "model_id": response_model},
|
||||
)
|
||||
raise ValueError("Model pricing not defined")
|
||||
|
||||
try:
|
||||
mspp = float(model_obj.sats_pricing.prompt)
|
||||
mspc = float(model_obj.sats_pricing.completion)
|
||||
mscr = float(model_obj.sats_pricing.input_cache_read or 0)
|
||||
mscw = float(model_obj.sats_pricing.input_cache_write or 0)
|
||||
|
||||
mspp_1k = mspp * 1_000_000.0
|
||||
mspc_1k = mspc * 1_000_000.0
|
||||
mscr_1k = mscr * 1_000_000.0 if mscr > 0 else mspp_1k
|
||||
mscw_1k = mscw * 1_000_000.0 if mscw > 0 else mspp_1k
|
||||
|
||||
logger.info(
|
||||
"Applied model-specific pricing",
|
||||
extra={
|
||||
"model": response_model,
|
||||
"input_price_msats_per_1k": mspp_1k,
|
||||
"output_price_msats_per_1k": mspc_1k,
|
||||
"cache_read_price_msats_per_1k": mscr_1k,
|
||||
"cache_write_price_msats_per_1k": mscw_1k,
|
||||
},
|
||||
)
|
||||
return mspp_1k, mspc_1k, mscr_1k, mscw_1k
|
||||
except Exception as e:
|
||||
logger.error("Invalid pricing data", extra={"error": str(e)})
|
||||
raise ValueError("Invalid pricing data") from e
|
||||
|
||||
|
||||
def _resolve_provider_fee(model_id: str) -> float:
|
||||
"""Resolve the provider fee multiplier for the given model id.
|
||||
|
||||
Falls back to 1.0 (no markup) when the provider cannot be resolved so
|
||||
the USD cost path never silently double-applies or omits the fee.
|
||||
"""
|
||||
from ..proxy import get_provider_for_model
|
||||
|
||||
if not model_id:
|
||||
return 1.0
|
||||
providers = get_provider_for_model(model_id)
|
||||
if not providers:
|
||||
return 1.0
|
||||
return float(providers[0].provider_fee)
|
||||
|
||||
|
||||
def _calculate_from_usd_cost(
|
||||
usd_cost: float,
|
||||
input_usd: float,
|
||||
output_usd: float,
|
||||
input_tokens: int,
|
||||
cache_read_tokens: int,
|
||||
cache_creation_tokens: int,
|
||||
output_tokens: int,
|
||||
response_data: dict,
|
||||
) -> CostData:
|
||||
"""Calculate cost from USD figures, deriving input/output split from tokens."""
|
||||
provider_fee = _resolve_provider_fee(response_data.get("model", ""))
|
||||
usd_cost = usd_cost * provider_fee
|
||||
input_usd = input_usd * provider_fee
|
||||
output_usd = output_usd * provider_fee
|
||||
sats_per_usd = 1.0 / sats_usd_price()
|
||||
cost_in_sats = usd_cost * sats_per_usd
|
||||
cost_in_msats = math.ceil(cost_in_sats * 1000)
|
||||
|
||||
if input_usd > 0 or output_usd > 0:
|
||||
input_msats = int((input_usd * sats_per_usd) * 1000)
|
||||
output_msats = int((output_usd * sats_per_usd) * 1000)
|
||||
else:
|
||||
effective_input_tokens = (
|
||||
input_tokens + cache_read_tokens + cache_creation_tokens
|
||||
)
|
||||
total_tokens = effective_input_tokens + output_tokens
|
||||
input_msats = (
|
||||
int(cost_in_msats * effective_input_tokens / total_tokens)
|
||||
if total_tokens > 0
|
||||
else 0
|
||||
)
|
||||
output_msats = cost_in_msats - input_msats
|
||||
|
||||
logger.info(
|
||||
"Calculated token-based cost",
|
||||
"Using cost from usage data/details",
|
||||
extra={
|
||||
"input_tokens": input_tokens,
|
||||
"output_tokens": output_tokens,
|
||||
"input_cost_msats": input_msats,
|
||||
"output_cost_msats": output_msats,
|
||||
"total_cost_msats": token_based_cost,
|
||||
"usd_cost": usd_cost,
|
||||
"cost_in_sats": cost_in_sats,
|
||||
"cost_in_msats": cost_in_msats,
|
||||
"model": response_data.get("model", "unknown"),
|
||||
},
|
||||
)
|
||||
|
||||
return CostData(
|
||||
base_msats=0,
|
||||
input_msats=int(input_msats),
|
||||
output_msats=int(output_msats),
|
||||
total_msats=token_based_cost,
|
||||
input_msats=input_msats,
|
||||
output_msats=output_msats,
|
||||
total_msats=cost_in_msats,
|
||||
total_usd=usd_cost,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cache_read_input_tokens=cache_read_tokens,
|
||||
cache_creation_input_tokens=cache_creation_tokens,
|
||||
cache_read_msats=0,
|
||||
cache_creation_msats=0,
|
||||
)
|
||||
|
||||
|
||||
def _calculate_from_tokens(
|
||||
input_tokens: int,
|
||||
output_tokens: int,
|
||||
cache_read_tokens: int,
|
||||
cache_creation_tokens: int,
|
||||
input_rate: float,
|
||||
output_rate: float,
|
||||
cache_read_rate: float,
|
||||
cache_creation_rate: float,
|
||||
response_data: dict,
|
||||
) -> CostData:
|
||||
"""Calculate cost from token counts using pricing rates."""
|
||||
calc_input_msats = round(input_tokens / 1000 * input_rate, 3)
|
||||
calc_output_msats = round(output_tokens / 1000 * output_rate, 3)
|
||||
calc_cache_read_msats = round(cache_read_tokens / 1000 * cache_read_rate, 3)
|
||||
calc_cache_write_msats = round(
|
||||
cache_creation_tokens / 1000 * cache_creation_rate, 3
|
||||
)
|
||||
token_based_cost = math.ceil(
|
||||
calc_input_msats
|
||||
+ calc_output_msats
|
||||
+ calc_cache_read_msats
|
||||
+ calc_cache_write_msats
|
||||
)
|
||||
total_usd = (token_based_cost / 1000.0) * sats_usd_price()
|
||||
|
||||
logger.info(
|
||||
"Calculated token-based cost",
|
||||
extra={
|
||||
"input_tokens": input_tokens,
|
||||
"output_tokens": output_tokens,
|
||||
"cache_read_input_tokens": cache_read_tokens,
|
||||
"cache_creation_input_tokens": cache_creation_tokens,
|
||||
"input_cost_msats": calc_input_msats,
|
||||
"output_cost_msats": calc_output_msats,
|
||||
"cache_read_cost_msats": calc_cache_read_msats,
|
||||
"cache_creation_cost_msats": calc_cache_write_msats,
|
||||
"total_cost_msats": token_based_cost,
|
||||
"total_usd": total_usd,
|
||||
"model": response_data.get("model", "unknown"),
|
||||
},
|
||||
)
|
||||
|
||||
return CostData(
|
||||
base_msats=0,
|
||||
input_msats=int(calc_input_msats),
|
||||
output_msats=int(calc_output_msats),
|
||||
total_msats=token_based_cost,
|
||||
total_usd=total_usd,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
cache_read_input_tokens=cache_read_tokens,
|
||||
cache_creation_input_tokens=cache_creation_tokens,
|
||||
cache_read_msats=int(calc_cache_read_msats),
|
||||
cache_creation_msats=int(calc_cache_write_msats),
|
||||
)
|
||||
|
||||
@@ -29,15 +29,13 @@ def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> N
|
||||
},
|
||||
)
|
||||
elif auth := headers.get("authorization", None):
|
||||
cashu_token = auth.split(" ")[1] if len(auth.split(" ")) > 1 else ""
|
||||
logger.debug(
|
||||
"Using Authorization header token",
|
||||
"Skipping preflight token balance check for Authorization header",
|
||||
extra={
|
||||
"token_preview": cashu_token[:20] + "..."
|
||||
if len(cashu_token) > 20
|
||||
else cashu_token
|
||||
"auth_preview": auth[:20] + "..." if len(auth) > 20 else auth,
|
||||
},
|
||||
)
|
||||
return
|
||||
else:
|
||||
logger.error("No authentication token provided")
|
||||
raise HTTPException(status_code=401, detail="Unauthorized")
|
||||
@@ -75,7 +73,7 @@ def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> N
|
||||
|
||||
if max_cost_for_model > amount_msat:
|
||||
raise HTTPException(
|
||||
status_code=413,
|
||||
status_code=402,
|
||||
detail={
|
||||
"reason": "Insufficient balance",
|
||||
"amount_required_msat": max_cost_for_model,
|
||||
@@ -169,9 +167,32 @@ async def calculate_discounted_max_cost(
|
||||
|
||||
tol = settings.tolerance_percentage
|
||||
tol_factor = max(0.0, 1 - float(tol) / 100.0)
|
||||
|
||||
max_prompt_allowed_sats = model_pricing.max_prompt_cost * tol_factor
|
||||
max_completion_allowed_sats = model_pricing.max_completion_cost * tol_factor
|
||||
|
||||
if model_obj:
|
||||
prompt_token_limit: int | None = None
|
||||
if model_obj.top_provider and (
|
||||
model_obj.top_provider.context_length
|
||||
or model_obj.top_provider.max_completion_tokens
|
||||
):
|
||||
cl = model_obj.top_provider.context_length
|
||||
mct = model_obj.top_provider.max_completion_tokens
|
||||
if cl and mct:
|
||||
prompt_token_limit = max(0, cl - mct)
|
||||
elif cl:
|
||||
prompt_token_limit = cl
|
||||
elif mct:
|
||||
prompt_token_limit = 0
|
||||
elif model_obj.context_length:
|
||||
prompt_token_limit = model_obj.context_length
|
||||
|
||||
if prompt_token_limit is not None:
|
||||
max_prompt_allowed_sats = (
|
||||
prompt_token_limit * model_pricing.prompt * tol_factor
|
||||
)
|
||||
|
||||
adjusted = max_cost_for_model
|
||||
|
||||
if messages := body.get("messages"):
|
||||
|
||||
@@ -25,82 +25,6 @@ class LNURLError(Exception):
|
||||
"""LNURL related errors."""
|
||||
|
||||
|
||||
def parse_lightning_invoice_amount(invoice: str, currency: str = "sat") -> int:
|
||||
"""Parse Lightning invoice (BOLT-11) to extract amount in specified currency units.
|
||||
|
||||
Args:
|
||||
invoice: BOLT-11 Lightning invoice string
|
||||
currency: Target currency unit ("sat" or "msat")
|
||||
|
||||
Returns:
|
||||
Amount in the specified currency unit
|
||||
|
||||
Raises:
|
||||
LNURLError: If invoice format is invalid or amount cannot be parsed
|
||||
"""
|
||||
invoice = invoice.lower().strip()
|
||||
|
||||
if not invoice.startswith("ln"):
|
||||
raise LNURLError("Invalid Lightning invoice format")
|
||||
|
||||
# Find the network part (bc, tb, etc.)
|
||||
network_start = 2
|
||||
while network_start < len(invoice) and invoice[network_start] not in "0123456789":
|
||||
network_start += 1
|
||||
|
||||
if network_start >= len(invoice):
|
||||
raise LNURLError("Invalid Lightning invoice format")
|
||||
|
||||
# Parse amount and multiplier
|
||||
amount_str = ""
|
||||
multiplier = ""
|
||||
i = network_start
|
||||
|
||||
# Extract numeric part
|
||||
while i < len(invoice) and invoice[i].isdigit():
|
||||
amount_str += invoice[i]
|
||||
i += 1
|
||||
|
||||
# Extract multiplier if present
|
||||
if i < len(invoice) and invoice[i] in "munp":
|
||||
multiplier = invoice[i]
|
||||
i += 1
|
||||
|
||||
# Check if we have the required "1" separator
|
||||
if i >= len(invoice) or invoice[i] != "1":
|
||||
raise LNURLError("Invalid Lightning invoice format")
|
||||
|
||||
if not amount_str:
|
||||
raise LNURLError("Lightning invoice amount not specified")
|
||||
|
||||
# Convert to base units
|
||||
try:
|
||||
amount = int(amount_str)
|
||||
except ValueError:
|
||||
raise LNURLError("Invalid Lightning invoice amount")
|
||||
|
||||
# Apply multiplier to get millisatoshis
|
||||
if multiplier == "m": # milli = 10^-3
|
||||
amount_msat = amount * 100_000_000 # amount is in BTC * 10^-3
|
||||
elif multiplier == "u": # micro = 10^-6
|
||||
amount_msat = amount * 100_000 # amount is in BTC * 10^-6
|
||||
elif multiplier == "n": # nano = 10^-9
|
||||
amount_msat = amount * 100 # amount is in BTC * 10^-9
|
||||
elif multiplier == "p": # pico = 10^-12
|
||||
amount_msat = amount // 10 # amount is in BTC * 10^-12
|
||||
else:
|
||||
# No multiplier means the amount is in BTC
|
||||
amount_msat = amount * 100_000_000_000 # Convert BTC to msat
|
||||
|
||||
# Convert to target currency unit
|
||||
if currency == "msat":
|
||||
return amount_msat
|
||||
elif currency == "sat":
|
||||
return amount_msat // 1000
|
||||
else:
|
||||
raise LNURLError(f"Unsupported currency for Lightning: {currency}")
|
||||
|
||||
|
||||
async def decode_lnurl(lnurl: str) -> str:
|
||||
"""Decode LNURL to get the actual URL.
|
||||
|
||||
@@ -291,9 +215,7 @@ async def raw_send_to_lnurl(
|
||||
lnurl_data["callback_url"], final_amount
|
||||
)
|
||||
|
||||
melt_quote_resp = await wallet.melt_quote(
|
||||
invoice=bolt11_invoice, amount_msat=final_amount
|
||||
)
|
||||
melt_quote_resp = await wallet.melt_quote(invoice=bolt11_invoice)
|
||||
|
||||
if amount:
|
||||
proofs, _ = await wallet.select_to_send(proofs, amount, set_reserved=True)
|
||||
|
||||
@@ -4,10 +4,11 @@ import random
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends
|
||||
from pydantic import BaseModel as V2BaseModel
|
||||
from pydantic.v1 import BaseModel
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from ..core.db import ModelRow, get_session
|
||||
from ..core.db import ModelRow, UpstreamProviderRow, get_session
|
||||
from ..core.logging import get_logger
|
||||
from ..core.settings import settings
|
||||
from .price import sats_usd_price
|
||||
@@ -60,6 +61,7 @@ class Model(BaseModel):
|
||||
upstream_provider_id: int | str | None = None
|
||||
canonical_slug: str | None = None
|
||||
alias_ids: list[str] | None = None
|
||||
forwarded_model_id: str | None = None
|
||||
|
||||
def __hash__(self) -> int:
|
||||
return hash(self.id)
|
||||
@@ -143,14 +145,6 @@ async def async_fetch_openrouter_models(source_filter: str | None = None) -> lis
|
||||
return []
|
||||
|
||||
|
||||
def is_openrouter_upstream() -> bool:
|
||||
try:
|
||||
base = (settings.upstream_base_url or "").strip().rstrip("/")
|
||||
except Exception:
|
||||
return False
|
||||
return base.lower() == "https://openrouter.ai/api/v1"
|
||||
|
||||
|
||||
def _row_to_model(
|
||||
row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01
|
||||
) -> Model:
|
||||
@@ -185,6 +179,7 @@ def _row_to_model(
|
||||
upstream_provider_id=row.upstream_provider_id,
|
||||
canonical_slug=getattr(row, "canonical_slug", None),
|
||||
alias_ids=json.loads(row.alias_ids) if row.alias_ids else None,
|
||||
forwarded_model_id=getattr(row, "forwarded_model_id", None) or row.id,
|
||||
)
|
||||
|
||||
if apply_provider_fee:
|
||||
@@ -203,33 +198,11 @@ def _row_to_model(
|
||||
return model
|
||||
|
||||
|
||||
def _model_to_row_payload(model: Model) -> dict[str, str | int | bool | None]:
|
||||
return {
|
||||
"id": model.id,
|
||||
"name": model.name,
|
||||
"created": model.created,
|
||||
"description": model.description,
|
||||
"context_length": model.context_length,
|
||||
"architecture": json.dumps(model.architecture.dict()),
|
||||
"pricing": json.dumps(model.pricing.dict()),
|
||||
"sats_pricing": json.dumps(model.sats_pricing.dict())
|
||||
if model.sats_pricing
|
||||
else None,
|
||||
"per_request_limits": json.dumps(model.per_request_limits)
|
||||
if model.per_request_limits is not None
|
||||
else None,
|
||||
"top_provider": json.dumps(model.top_provider.dict())
|
||||
if model.top_provider is not None
|
||||
else None,
|
||||
"enabled": model.enabled,
|
||||
"upstream_provider_id": model.upstream_provider_id,
|
||||
}
|
||||
|
||||
|
||||
async def list_models(
|
||||
session: AsyncSession,
|
||||
upstream_id: int,
|
||||
include_disabled: bool = False,
|
||||
apply_fees: bool = True,
|
||||
) -> list[Model]:
|
||||
from sqlmodel import select
|
||||
|
||||
@@ -247,7 +220,7 @@ async def list_models(
|
||||
return [
|
||||
_row_to_model(
|
||||
r,
|
||||
apply_provider_fee=True,
|
||||
apply_provider_fee=apply_fees,
|
||||
provider_fee=providers_by_id[r.upstream_provider_id].provider_fee
|
||||
if r.upstream_provider_id in providers_by_id
|
||||
else 1.01,
|
||||
@@ -261,21 +234,6 @@ async def list_models(
|
||||
]
|
||||
|
||||
|
||||
async def get_model_by_id(
|
||||
model_id: str, provider_id: int, session: AsyncSession
|
||||
) -> Model | None:
|
||||
from ..core.db import UpstreamProviderRow
|
||||
|
||||
row = await session.get(ModelRow, (model_id, provider_id))
|
||||
if not row or not row.enabled:
|
||||
return None
|
||||
provider = await session.get(UpstreamProviderRow, provider_id)
|
||||
if not provider or not provider.enabled:
|
||||
return None
|
||||
provider_fee = provider.provider_fee if provider else 1.01
|
||||
return _row_to_model(row, apply_provider_fee=True, provider_fee=provider_fee)
|
||||
|
||||
|
||||
def _calculate_usd_max_costs(model: Model) -> tuple[float, float, float]:
|
||||
"""Calculate max costs in USD based on model context/token limits.
|
||||
|
||||
@@ -374,6 +332,7 @@ def _update_model_sats_pricing(model: Model, sats_to_usd: float) -> Model:
|
||||
upstream_provider_id=model.upstream_provider_id,
|
||||
canonical_slug=model.canonical_slug,
|
||||
alias_ids=model.alias_ids,
|
||||
forwarded_model_id=model.forwarded_model_id,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
@@ -392,6 +351,9 @@ async def _update_sats_pricing_once() -> None:
|
||||
from ..proxy import get_upstreams, refresh_model_maps
|
||||
|
||||
upstreams = get_upstreams()
|
||||
if not upstreams:
|
||||
return
|
||||
|
||||
sats_to_usd = sats_usd_price()
|
||||
|
||||
updated_count = 0
|
||||
@@ -401,11 +363,14 @@ async def _update_sats_pricing_once() -> None:
|
||||
for m in upstream.get_cached_models()
|
||||
]
|
||||
upstream._models_cache = updated_models
|
||||
upstream._models_by_id = {m.id: m for m in updated_models}
|
||||
upstream._models_by_id = {m.forwarded_model_id or m.id: m for m in updated_models}
|
||||
updated_count += len(updated_models)
|
||||
|
||||
if updated_count > 0:
|
||||
logger.info("Updated sats pricing", extra={"models_updated": updated_count})
|
||||
logger.info(
|
||||
f"Updated sats pricing for {updated_count} models",
|
||||
extra={"models_updated": updated_count},
|
||||
)
|
||||
await refresh_model_maps()
|
||||
|
||||
|
||||
@@ -447,11 +412,89 @@ async def update_sats_pricing() -> None:
|
||||
logger.error(f"Error updating sats pricing: {e}")
|
||||
|
||||
|
||||
class ModelTestRequest(V2BaseModel):
|
||||
model_id: str
|
||||
endpoint_type: str
|
||||
request_data: dict
|
||||
|
||||
|
||||
@models_router.post("/api/models/test")
|
||||
async def test_model(
|
||||
payload: ModelTestRequest,
|
||||
session: AsyncSession = Depends(get_session),
|
||||
) -> dict:
|
||||
"""Test a model by sending a request through its configured upstream provider."""
|
||||
from sqlmodel import select
|
||||
|
||||
result = await session.execute(
|
||||
select(ModelRow).where(ModelRow.id == payload.model_id)
|
||||
)
|
||||
model_row = result.scalars().first()
|
||||
|
||||
if not model_row:
|
||||
return {
|
||||
"success": False,
|
||||
"error": f"Model '{payload.model_id}' not found in database",
|
||||
"status_code": 404,
|
||||
}
|
||||
|
||||
provider = await session.get(UpstreamProviderRow, model_row.upstream_provider_id)
|
||||
if not provider:
|
||||
return {
|
||||
"success": False,
|
||||
"error": "Upstream provider not found",
|
||||
"status_code": 404,
|
||||
}
|
||||
|
||||
base_url = provider.base_url.rstrip("/")
|
||||
if payload.endpoint_type == "chat-completions":
|
||||
url = f"{base_url}/chat/completions"
|
||||
else:
|
||||
url = f"{base_url}/{payload.endpoint_type}"
|
||||
|
||||
actual_model_id = model_row.forwarded_model_id or model_row.id
|
||||
request_data = dict(payload.request_data)
|
||||
request_data["model"] = actual_model_id
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {provider.api_key}",
|
||||
}
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=30.0) as client:
|
||||
response = await client.post(url, json=request_data, headers=headers)
|
||||
try:
|
||||
response_data = response.json()
|
||||
except Exception:
|
||||
response_data = {"raw": response.text}
|
||||
|
||||
return {
|
||||
"success": response.status_code < 400,
|
||||
"data": response_data,
|
||||
"status_code": response.status_code,
|
||||
}
|
||||
except Exception as e:
|
||||
return {
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"status_code": 500,
|
||||
}
|
||||
|
||||
|
||||
@models_router.get("/v1/models")
|
||||
@models_router.get("/models", include_in_schema=False)
|
||||
@models_router.get("/v1/models/", include_in_schema=False)
|
||||
@models_router.get("/models")
|
||||
@models_router.get("/models/", include_in_schema=False)
|
||||
async def models(session: AsyncSession = Depends(get_session)) -> dict:
|
||||
"""Get all available models from all providers with database overrides applied."""
|
||||
from ..proxy import get_unique_models
|
||||
|
||||
items = get_unique_models()
|
||||
return {"data": items}
|
||||
data = []
|
||||
for model in items:
|
||||
m = model.dict()
|
||||
if model.forwarded_model_id:
|
||||
m["id"] = model.forwarded_model_id
|
||||
data.append(m)
|
||||
return {"data": data}
|
||||
|
||||
398
routstr/proxy.py
398
routstr/proxy.py
@@ -16,6 +16,9 @@ from .core.db import (
|
||||
create_session,
|
||||
get_session,
|
||||
)
|
||||
from .core.exceptions import UpstreamError
|
||||
from .core.not_found import build_not_found_response
|
||||
from .core.settings import settings
|
||||
from .payment.helpers import (
|
||||
calculate_discounted_max_cost,
|
||||
check_token_balance,
|
||||
@@ -31,7 +34,9 @@ proxy_router = APIRouter()
|
||||
|
||||
_upstreams: list[BaseUpstreamProvider] = []
|
||||
_model_instances: dict[str, Model] = {} # All aliases -> Model
|
||||
_provider_map: dict[str, BaseUpstreamProvider] = {} # All aliases -> Provider
|
||||
_provider_map: dict[
|
||||
str, list[BaseUpstreamProvider]
|
||||
] = {} # All aliases -> List[Provider]
|
||||
_unique_models: dict[str, Model] = {} # Unique model.id -> Model (no duplicates)
|
||||
|
||||
|
||||
@@ -65,11 +70,29 @@ def get_upstreams() -> list[BaseUpstreamProvider]:
|
||||
|
||||
def get_model_instance(model_id: str) -> Model | None:
|
||||
"""Get Model instance by ID from global cache."""
|
||||
return _model_instances.get(model_id.lower())
|
||||
if not model_id:
|
||||
return None
|
||||
|
||||
model_id_lower = model_id.lower()
|
||||
# Try exact match first
|
||||
if model := _model_instances.get(model_id_lower):
|
||||
return model
|
||||
|
||||
# Try stripping common version suffixes (e.g., -20251222)
|
||||
# This handles cases where upstream returns a specific version
|
||||
# but we only track the base model name.
|
||||
import re
|
||||
|
||||
base_model_id = re.sub(r"-\d{8}$", "", model_id_lower)
|
||||
if base_model_id != model_id_lower:
|
||||
if model := _model_instances.get(base_model_id):
|
||||
return model
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def get_provider_for_model(model_id: str) -> BaseUpstreamProvider | None:
|
||||
"""Get UpstreamProvider for model ID from global cache."""
|
||||
def get_provider_for_model(model_id: str) -> list[BaseUpstreamProvider] | None:
|
||||
"""Get UpstreamProvider list for model ID from global cache."""
|
||||
return _provider_map.get(model_id.lower())
|
||||
|
||||
|
||||
@@ -128,34 +151,86 @@ async def refresh_model_maps_periodically() -> None:
|
||||
)
|
||||
|
||||
|
||||
_API_PATH_PREFIXES = (
|
||||
"v1/",
|
||||
"responses",
|
||||
"chat/",
|
||||
"completions",
|
||||
"models",
|
||||
"embeddings",
|
||||
"audio/",
|
||||
"images/",
|
||||
"moderations",
|
||||
"providers",
|
||||
"tee/",
|
||||
)
|
||||
|
||||
|
||||
@proxy_router.api_route("/{path:path}", methods=["GET", "POST"], response_model=None)
|
||||
async def proxy(
|
||||
request: Request, path: str, session: AsyncSession = Depends(get_session)
|
||||
) -> Response | StreamingResponse:
|
||||
headers = dict(request.headers)
|
||||
# GET requests must hit a known API prefix; otherwise return a 404 (HTML
|
||||
# for browsers, JSON for API clients). POST requests are always forwarded
|
||||
# so that OpenAI-style endpoints work with or without the `v1/` prefix
|
||||
# (e.g. `/chat/completions` as well as `/v1/chat/completions`).
|
||||
if request.method == "GET" and not path.startswith(_API_PATH_PREFIXES):
|
||||
return build_not_found_response(request, path)
|
||||
|
||||
if "x-cashu" not in headers and "authorization" not in headers.keys():
|
||||
return create_error_response(
|
||||
"unauthorized", "Unauthorized", 401, request=request
|
||||
)
|
||||
headers = dict(request.headers)
|
||||
|
||||
is_responses_api = path.startswith("v1/responses") or path.startswith("responses")
|
||||
request_body = await request.body()
|
||||
request_body_dict = parse_request_body_json(request_body, path)
|
||||
|
||||
# /tee/* GET requests (e.g. attestation) don't map to models — just
|
||||
# forward to all enabled upstreams without model/cost/auth lookups.
|
||||
if request.method == "GET" and path.startswith("tee/"):
|
||||
all_upstreams = _upstreams
|
||||
last_error_response = None
|
||||
for i, upstream in enumerate(all_upstreams):
|
||||
try:
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
response = await upstream.forward_get_request(request, path, headers)
|
||||
if response.status_code in [502, 429] and i < len(all_upstreams) - 1:
|
||||
logger.warning(
|
||||
"Upstream %s returned %s for tee GET %s, trying next",
|
||||
upstream.provider_type,
|
||||
response.status_code,
|
||||
path,
|
||||
)
|
||||
continue
|
||||
return response
|
||||
except UpstreamError as e:
|
||||
logger.warning(
|
||||
"Upstream %s failed for tee GET %s: %s",
|
||||
upstream.provider_type,
|
||||
path,
|
||||
e,
|
||||
)
|
||||
if i == len(all_upstreams) - 1:
|
||||
last_error_response = create_error_response(
|
||||
"upstream_error", str(e), 502, request=request
|
||||
)
|
||||
continue
|
||||
return last_error_response or create_error_response(
|
||||
"upstream_error", "All upstreams failed", 502, request=request
|
||||
)
|
||||
|
||||
if is_responses_api:
|
||||
model_id = extract_model_from_responses_request(request_body_dict)
|
||||
else:
|
||||
model_id = request_body_dict.get("model", "unknown")
|
||||
|
||||
model_obj = get_model_instance(model_id)
|
||||
|
||||
if not model_obj:
|
||||
return create_error_response(
|
||||
"invalid_model", f"Model '{model_id}' not found", 400, request=request
|
||||
)
|
||||
|
||||
upstream = get_provider_for_model(model_id)
|
||||
if not upstream:
|
||||
upstreams = get_provider_for_model(model_id)
|
||||
if not upstreams:
|
||||
return create_error_response(
|
||||
"invalid_model",
|
||||
f"No provider found for model '{model_id}'",
|
||||
@@ -163,26 +238,60 @@ async def proxy(
|
||||
request=request,
|
||||
)
|
||||
|
||||
# todo figure out cost calculation since fallback provider is usually not the same price
|
||||
# Use first provider for initial checks/cost calculation
|
||||
# primary_upstream = upstreams[0]
|
||||
|
||||
_max_cost_for_model = await get_max_cost_for_model(
|
||||
model=model_id, session=session, model_obj=model_obj
|
||||
)
|
||||
max_cost_for_model = await calculate_discounted_max_cost(
|
||||
_max_cost_for_model, request_body_dict, model_obj=model_obj
|
||||
)
|
||||
# Ensure max_cost_for_model is at least the minimum allowed request cost
|
||||
max_cost_for_model = max(max_cost_for_model, settings.min_request_msat)
|
||||
|
||||
check_token_balance(headers, request_body_dict, max_cost_for_model)
|
||||
|
||||
if x_cashu := headers.get("x-cashu", None):
|
||||
if is_responses_api:
|
||||
return await upstream.handle_x_cashu_responses(
|
||||
request, x_cashu, path, max_cost_for_model, model_obj
|
||||
)
|
||||
else:
|
||||
return await upstream.handle_x_cashu(
|
||||
request, x_cashu, path, max_cost_for_model, model_obj
|
||||
)
|
||||
last_error = None
|
||||
for i, upstream in enumerate(upstreams):
|
||||
try:
|
||||
if is_responses_api:
|
||||
return await upstream.handle_x_cashu_responses(
|
||||
request, x_cashu, path, max_cost_for_model, model_obj
|
||||
)
|
||||
else:
|
||||
return await upstream.handle_x_cashu(
|
||||
request, x_cashu, path, max_cost_for_model, model_obj
|
||||
)
|
||||
except UpstreamError as e:
|
||||
logger.warning(
|
||||
"Upstream %s failed (x-cashu) for model=%s: %s",
|
||||
upstream.provider_type,
|
||||
model_id,
|
||||
e,
|
||||
extra={
|
||||
"provider": upstream.provider_type,
|
||||
"model": model_id,
|
||||
"status_code": e.status_code,
|
||||
},
|
||||
)
|
||||
if i == len(upstreams) - 1:
|
||||
last_error = e
|
||||
continue
|
||||
|
||||
return create_error_response(
|
||||
"upstream_error",
|
||||
str(last_error) if last_error else "All upstreams failed",
|
||||
502,
|
||||
request=request,
|
||||
)
|
||||
|
||||
elif auth := headers.get("authorization", None):
|
||||
key = await get_bearer_token_key(headers, path, session, auth)
|
||||
key = await get_bearer_token_key(
|
||||
headers, path, session, auth, max_cost_for_model, model_id
|
||||
)
|
||||
|
||||
else:
|
||||
if request.method not in ["GET"]:
|
||||
@@ -194,77 +303,199 @@ async def proxy(
|
||||
)
|
||||
|
||||
logger.debug("Processing unauthenticated GET request", extra={"path": path})
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
return await upstream.forward_get_request(request, path, headers)
|
||||
|
||||
last_error_response = None
|
||||
for i, upstream in enumerate(upstreams):
|
||||
try:
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
response = await upstream.forward_get_request(request, path, headers)
|
||||
|
||||
if response.status_code in [502, 429] and i < len(upstreams) - 1:
|
||||
error_message = ""
|
||||
try:
|
||||
if hasattr(response, "body"):
|
||||
body_bytes = response.body
|
||||
data = json.loads(body_bytes)
|
||||
if "error" in data:
|
||||
error_data = data["error"]
|
||||
if isinstance(error_data, dict):
|
||||
error_message = error_data.get("message", "")
|
||||
elif isinstance(error_data, str):
|
||||
error_message = error_data
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
await upstream.on_upstream_error_redirect(
|
||||
response.status_code, error_message
|
||||
)
|
||||
|
||||
logger.warning(
|
||||
f"Upstream {upstream.provider_type} returned {response.status_code} (GET), trying next provider",
|
||||
extra={
|
||||
"status_code": response.status_code,
|
||||
"upstream": upstream.provider_type,
|
||||
},
|
||||
)
|
||||
continue
|
||||
return response
|
||||
except UpstreamError as e:
|
||||
logger.warning(f"Upstream {upstream.provider_type} failed (GET): {e}")
|
||||
if i == len(upstreams) - 1:
|
||||
last_error_response = create_error_response(
|
||||
"upstream_error", str(e), 502, request=request
|
||||
)
|
||||
continue
|
||||
return last_error_response or create_error_response(
|
||||
"upstream_error", "All upstreams failed", 502, request=request
|
||||
)
|
||||
|
||||
if request_body_dict:
|
||||
await pay_for_request(key, max_cost_for_model, session)
|
||||
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
for i, upstream in enumerate(upstreams):
|
||||
headers = upstream.prepare_headers(dict(request.headers))
|
||||
|
||||
try:
|
||||
if is_responses_api:
|
||||
response = await upstream.forward_responses_request(
|
||||
request,
|
||||
path,
|
||||
headers,
|
||||
request_body,
|
||||
key,
|
||||
max_cost_for_model,
|
||||
session,
|
||||
model_obj,
|
||||
try:
|
||||
try:
|
||||
if is_responses_api:
|
||||
response = await upstream.forward_responses_request(
|
||||
request,
|
||||
path,
|
||||
headers,
|
||||
request_body,
|
||||
key,
|
||||
max_cost_for_model,
|
||||
session,
|
||||
model_obj,
|
||||
)
|
||||
else:
|
||||
response = await upstream.forward_request(
|
||||
request,
|
||||
path,
|
||||
headers,
|
||||
request_body,
|
||||
key,
|
||||
max_cost_for_model,
|
||||
session,
|
||||
model_obj,
|
||||
)
|
||||
except UpstreamError:
|
||||
# Let the outer UpstreamError handler manage retry/revert
|
||||
raise
|
||||
except Exception as e:
|
||||
# Unexpected error (not an upstream failure) — revert and propagate
|
||||
logger.error(
|
||||
"Unexpected error in upstream request, reverting payment",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
"path": path,
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
"max_cost_for_model": max_cost_for_model,
|
||||
},
|
||||
)
|
||||
await revert_pay_for_request(key, session, max_cost_for_model)
|
||||
raise
|
||||
|
||||
if response.status_code != 200:
|
||||
# Check if we should retry (502 Upstream Error or 429 Rate Limit)
|
||||
should_retry = response.status_code in [502, 429, 400, 401, 403, 404]
|
||||
if should_retry and i < len(upstreams) - 1:
|
||||
error_message = ""
|
||||
try:
|
||||
if hasattr(response, "body"):
|
||||
body_bytes = response.body
|
||||
data = json.loads(body_bytes)
|
||||
if "error" in data:
|
||||
error_data = data["error"]
|
||||
if isinstance(error_data, dict):
|
||||
error_message = error_data.get("message", "")
|
||||
elif isinstance(error_data, str):
|
||||
error_message = error_data
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
await upstream.on_upstream_error_redirect(
|
||||
response.status_code, error_message
|
||||
)
|
||||
|
||||
logger.warning(
|
||||
"Upstream %s returned %s for model=%s, trying next provider",
|
||||
upstream.provider_type,
|
||||
response.status_code,
|
||||
model_id,
|
||||
extra={
|
||||
"status_code": response.status_code,
|
||||
"provider": upstream.provider_type,
|
||||
"model": model_id,
|
||||
},
|
||||
)
|
||||
continue
|
||||
|
||||
# 4xx error (user error), or other non-retryable error, or last provider failed
|
||||
await revert_pay_for_request(key, session, max_cost_for_model)
|
||||
logger.warning(
|
||||
"Upstream request failed, revert payment "
|
||||
"(provider=%s model=%s status=%s path=%s)",
|
||||
upstream.provider_type,
|
||||
model_id,
|
||||
response.status_code,
|
||||
path,
|
||||
extra={
|
||||
"status_code": response.status_code,
|
||||
"path": path,
|
||||
"provider": upstream.provider_type,
|
||||
"model": model_id,
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
"key_balance": key.balance,
|
||||
"max_cost_for_model": max_cost_for_model,
|
||||
},
|
||||
)
|
||||
return response
|
||||
|
||||
return response
|
||||
|
||||
except UpstreamError as e:
|
||||
logger.warning(
|
||||
"Upstream %s failed for model=%s: %s",
|
||||
upstream.provider_type,
|
||||
model_id,
|
||||
e,
|
||||
extra={
|
||||
"provider": upstream.provider_type,
|
||||
"model": model_id,
|
||||
"status_code": e.status_code,
|
||||
"retry": i < len(upstreams) - 1,
|
||||
},
|
||||
)
|
||||
else:
|
||||
response = await upstream.forward_request(
|
||||
request,
|
||||
path,
|
||||
headers,
|
||||
request_body,
|
||||
key,
|
||||
max_cost_for_model,
|
||||
session,
|
||||
model_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Upstream request failed, ensuring payment is reverted",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
"path": path,
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
"max_cost_for_model": max_cost_for_model,
|
||||
},
|
||||
)
|
||||
await revert_pay_for_request(key, session, max_cost_for_model)
|
||||
raise
|
||||
|
||||
if response.status_code != 200:
|
||||
await revert_pay_for_request(key, session, max_cost_for_model)
|
||||
logger.warning(
|
||||
"Upstream request failed, revert payment",
|
||||
extra={
|
||||
"status_code": response.status_code,
|
||||
"path": path,
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
"key_balance": key.balance,
|
||||
"max_cost_for_model": max_cost_for_model,
|
||||
"upstream_headers": response.headers
|
||||
if hasattr(response, "headers")
|
||||
else None,
|
||||
},
|
||||
)
|
||||
# Return the mapped error response generated earlier rather than masking with 502
|
||||
return response
|
||||
# If this was the last provider
|
||||
if i == len(upstreams) - 1:
|
||||
await revert_pay_for_request(key, session, max_cost_for_model)
|
||||
return create_error_response(
|
||||
"upstream_error", str(e), 502, request=request
|
||||
)
|
||||
|
||||
return response
|
||||
# Otherwise loop continues to next provider
|
||||
continue
|
||||
|
||||
# Should not be reached given logic above
|
||||
return create_error_response(
|
||||
"upstream_error", "All upstreams failed", 502, request=request
|
||||
)
|
||||
|
||||
|
||||
async def get_bearer_token_key(
|
||||
headers: dict, path: str, session: AsyncSession, auth: str
|
||||
headers: dict,
|
||||
path: str,
|
||||
session: AsyncSession,
|
||||
auth: str,
|
||||
min_cost: int = 0,
|
||||
model_id: str = "unknown",
|
||||
) -> ApiKey:
|
||||
"""Handle bearer token authentication proxy requests."""
|
||||
bearer_key = auth.replace("Bearer ", "") if auth.startswith("Bearer ") else ""
|
||||
parts = auth.split()
|
||||
bearer_key = parts[1] if len(parts) > 1 and parts[0].lower() == "bearer" else ""
|
||||
refund_address = headers.get("Refund-LNURL", None)
|
||||
key_expiry_time = headers.get("Key-Expiry-Time", None)
|
||||
|
||||
@@ -277,6 +508,7 @@ async def get_bearer_token_key(
|
||||
"bearer_key_preview": bearer_key[:20] + "..."
|
||||
if len(bearer_key) > 20
|
||||
else bearer_key,
|
||||
"min_cost": min_cost,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -315,6 +547,7 @@ async def get_bearer_token_key(
|
||||
session,
|
||||
refund_address,
|
||||
key_expiry_time, # type: ignore
|
||||
min_cost=min_cost,
|
||||
)
|
||||
logger.info(
|
||||
"Bearer token validated successfully",
|
||||
@@ -326,15 +559,16 @@ async def get_bearer_token_key(
|
||||
)
|
||||
return key
|
||||
except Exception as e:
|
||||
key_preview = bearer_key[:20] + "..." if len(bearer_key) > 20 else bearer_key
|
||||
logger.error(
|
||||
"Bearer token validation failed",
|
||||
f"Bearer token validation failed: {type(e).__name__}: {e} path={path} model={model_id!r} min_cost={min_cost} key={key_preview!r}",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
"path": path,
|
||||
"bearer_key_preview": bearer_key[:20] + "..."
|
||||
if len(bearer_key) > 20
|
||||
else bearer_key,
|
||||
"model_id": model_id,
|
||||
"min_cost_msat": min_cost,
|
||||
"bearer_key_preview": key_preview,
|
||||
},
|
||||
)
|
||||
raise
|
||||
|
||||
@@ -10,6 +10,7 @@ from .openai import OpenAIUpstreamProvider
|
||||
from .openrouter import OpenRouterUpstreamProvider
|
||||
from .perplexity import PerplexityUpstreamProvider
|
||||
from .ppqai import PPQAIUpstreamProvider
|
||||
from .routstr import RoutstrUpstreamProvider
|
||||
from .xai import XAIUpstreamProvider
|
||||
|
||||
upstream_provider_classes: list[type[BaseUpstreamProvider]] = [
|
||||
@@ -24,6 +25,7 @@ upstream_provider_classes: list[type[BaseUpstreamProvider]] = [
|
||||
OpenRouterUpstreamProvider,
|
||||
PerplexityUpstreamProvider,
|
||||
PPQAIUpstreamProvider,
|
||||
RoutstrUpstreamProvider,
|
||||
XAIUpstreamProvider,
|
||||
]
|
||||
"""List of all upstream classes"""
|
||||
|
||||
@@ -13,6 +13,8 @@ class AnthropicUpstreamProvider(BaseUpstreamProvider):
|
||||
provider_type = "anthropic"
|
||||
default_base_url = "https://api.anthropic.com/v1"
|
||||
platform_url = "https://console.anthropic.com/settings/keys"
|
||||
supports_anthropic_messages = True
|
||||
litellm_provider_prefix = "anthropic/"
|
||||
|
||||
def __init__(self, api_key: str, provider_fee: float = 1.01):
|
||||
super().__init__(
|
||||
|
||||
157
routstr/upstream/auto_topup.py
Normal file
157
routstr/upstream/auto_topup.py
Normal file
@@ -0,0 +1,157 @@
|
||||
import asyncio
|
||||
import json
|
||||
|
||||
from sqlmodel import select
|
||||
|
||||
from ..core import get_logger
|
||||
from ..core.db import UpstreamProviderRow, create_session
|
||||
from ..wallet import send_token
|
||||
from .routstr import RoutstrUpstreamProvider
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
# Check every 60 seconds
|
||||
AUTO_TOPUP_INTERVAL_SECONDS = 60
|
||||
|
||||
|
||||
async def periodic_auto_topup() -> None:
|
||||
"""Background task that monitors Routstr provider balances and auto-tops up when below threshold.
|
||||
|
||||
For each Routstr provider with auto_topup enabled in provider_settings:
|
||||
1. Checks the upstream balance via get_balance()
|
||||
2. If balance < topup_threshold, creates a cashu token from the configured mint
|
||||
3. Sends the token to the upstream provider via topup()
|
||||
"""
|
||||
# Wait for initial startup to complete
|
||||
await asyncio.sleep(30)
|
||||
logger.info("Auto top-up worker started")
|
||||
|
||||
while True:
|
||||
try:
|
||||
await _run_auto_topup_cycle()
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Auto top-up cycle failed",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
)
|
||||
|
||||
await asyncio.sleep(AUTO_TOPUP_INTERVAL_SECONDS)
|
||||
|
||||
|
||||
async def _run_auto_topup_cycle() -> None:
|
||||
"""Single cycle: check all eligible providers and top up if needed."""
|
||||
async with create_session() as session:
|
||||
query = select(UpstreamProviderRow).where(
|
||||
UpstreamProviderRow.provider_type == "routstr",
|
||||
UpstreamProviderRow.enabled == True, # noqa: E712
|
||||
)
|
||||
result = await session.exec(query)
|
||||
providers = result.all()
|
||||
|
||||
for row in providers:
|
||||
try:
|
||||
await _check_and_topup(row)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Auto top-up failed for provider",
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"base_url": row.base_url,
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def _check_and_topup(row: UpstreamProviderRow) -> None:
|
||||
"""Check a single provider's balance and top up if below threshold."""
|
||||
# Parse provider settings
|
||||
settings: dict = {}
|
||||
if row.provider_settings:
|
||||
try:
|
||||
settings = json.loads(row.provider_settings)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return
|
||||
|
||||
if not settings.get("auto_topup"):
|
||||
return
|
||||
|
||||
threshold = settings.get("topup_threshold")
|
||||
amount = settings.get("topup_amount_limit")
|
||||
mint_url = settings.get("topup_mint_url")
|
||||
|
||||
if not threshold or not amount or not mint_url:
|
||||
logger.warning(
|
||||
"Auto top-up enabled but missing configuration",
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"has_threshold": bool(threshold),
|
||||
"has_amount": bool(amount),
|
||||
"has_mint": bool(mint_url),
|
||||
},
|
||||
)
|
||||
return
|
||||
|
||||
if not row.api_key:
|
||||
return
|
||||
|
||||
# Instantiate provider and check balance
|
||||
provider = RoutstrUpstreamProvider.from_db_row(row)
|
||||
balance = await provider.get_balance()
|
||||
|
||||
if balance is None:
|
||||
logger.warning(
|
||||
"Could not fetch balance for auto top-up",
|
||||
extra={"provider_id": row.id, "base_url": row.base_url},
|
||||
)
|
||||
return
|
||||
|
||||
if balance >= threshold * 1000:
|
||||
return
|
||||
|
||||
# Balance is below threshold - create token and top up
|
||||
logger.info(
|
||||
"Auto top-up triggered",
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"balance": balance,
|
||||
"threshold": threshold,
|
||||
"topup_amount": amount,
|
||||
"mint_url": mint_url,
|
||||
},
|
||||
)
|
||||
|
||||
print(amount, mint_url)
|
||||
try:
|
||||
token = await send_token(amount, "sat", mint_url)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to create cashu token for auto top-up",
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"amount": amount,
|
||||
"mint_url": mint_url,
|
||||
"error": str(e),
|
||||
},
|
||||
)
|
||||
return
|
||||
|
||||
result = await provider.topup(token)
|
||||
|
||||
if "error" in result:
|
||||
logger.error(
|
||||
"Auto top-up upstream call failed",
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"error": result["error"],
|
||||
},
|
||||
)
|
||||
else:
|
||||
logger.info(
|
||||
"Auto top-up completed successfully",
|
||||
extra={
|
||||
"provider_id": row.id,
|
||||
"amount": amount,
|
||||
"new_balance_approx": balance + amount,
|
||||
},
|
||||
)
|
||||
@@ -4,6 +4,7 @@ from .base import BaseUpstreamProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import UpstreamProviderRow
|
||||
from ..payment.models import Model
|
||||
|
||||
|
||||
class AzureUpstreamProvider(BaseUpstreamProvider):
|
||||
@@ -12,6 +13,7 @@ class AzureUpstreamProvider(BaseUpstreamProvider):
|
||||
provider_type = "azure"
|
||||
default_base_url = None
|
||||
platform_url = "https://portal.azure.com/"
|
||||
litellm_provider_prefix = "azure/"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -58,19 +60,52 @@ class AzureUpstreamProvider(BaseUpstreamProvider):
|
||||
"platform_url": cls.platform_url,
|
||||
}
|
||||
|
||||
def prepare_headers(self, request_headers: dict) -> dict:
|
||||
"""Prepare headers for Azure OpenAI, adding api-key."""
|
||||
headers = super().prepare_headers(request_headers)
|
||||
if self.api_key:
|
||||
headers["api-key"] = self.api_key
|
||||
headers.pop("Authorization", None)
|
||||
headers.pop("authorization", None)
|
||||
return headers
|
||||
|
||||
def prepare_params(
|
||||
self, path: str, query_params: Mapping[str, str] | None
|
||||
) -> Mapping[str, str]:
|
||||
"""Prepare query parameters for Azure OpenAI, adding API version.
|
||||
|
||||
Args:
|
||||
path: Request path
|
||||
query_params: Original query parameters from the client
|
||||
|
||||
Returns:
|
||||
Query parameters dict with Azure API version added for chat completions
|
||||
"""
|
||||
"""Prepare query parameters for Azure OpenAI, adding API version."""
|
||||
params = dict(query_params or {})
|
||||
if path.endswith("chat/completions"):
|
||||
params["api-version"] = self.api_version
|
||||
version = (self.api_version or "").replace("\ufeff", "").strip()
|
||||
if not version or version.lower() == "v1":
|
||||
version = "2024-02-15-preview"
|
||||
params["api-version"] = version
|
||||
return params
|
||||
|
||||
def normalize_request_path(
|
||||
self, path: str, model_obj: "Model | None" = None
|
||||
) -> str:
|
||||
"""Build Azure deployment-specific request path."""
|
||||
clean_path = super().normalize_request_path(path, model_obj).lstrip("/")
|
||||
if model_obj is None:
|
||||
return clean_path
|
||||
|
||||
deployment_id = getattr(
|
||||
model_obj, "canonical_slug", None
|
||||
) or self.transform_model_name(model_obj.id)
|
||||
deployment_id = deployment_id.split("/")[-1]
|
||||
return f"openai/deployments/{deployment_id}/{clean_path}"
|
||||
|
||||
def get_request_base_url(
|
||||
self, path: str, model_obj: "Model | None" = None
|
||||
) -> str:
|
||||
"""Use endpoint root, stripping accidental /openai/v1 suffix if present."""
|
||||
base_url = self.base_url.rstrip("/")
|
||||
marker = "/openai/v1"
|
||||
if marker in base_url:
|
||||
base_url = base_url.split(marker, 1)[0].rstrip("/")
|
||||
return base_url
|
||||
|
||||
def transform_model_name(self, model_id: str) -> str:
|
||||
"""Extract deployment name from model ID."""
|
||||
if "/" in model_id:
|
||||
return model_id.split("/")[-1]
|
||||
return model_id
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -26,14 +26,14 @@ class GeminiClient(BaseAPIClient):
|
||||
max_tokens: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> dict[str, Any]:
|
||||
from openai import NOT_GIVEN
|
||||
from openai import omit
|
||||
|
||||
response = await self.client.chat.completions.create(
|
||||
model=model,
|
||||
messages=messages, # type: ignore
|
||||
temperature=temperature if temperature is not None else NOT_GIVEN,
|
||||
max_tokens=max_tokens if max_tokens is not None else NOT_GIVEN,
|
||||
top_p=kwargs.get("top_p", NOT_GIVEN),
|
||||
temperature=temperature if temperature is not None else omit,
|
||||
max_tokens=max_tokens if max_tokens is not None else omit,
|
||||
top_p=kwargs.get("top_p", omit),
|
||||
)
|
||||
return response.model_dump()
|
||||
|
||||
@@ -45,7 +45,7 @@ class GeminiClient(BaseAPIClient):
|
||||
max_tokens: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncGenerator[dict[str, Any], None]:
|
||||
from openai import NOT_GIVEN
|
||||
from openai import omit
|
||||
|
||||
usage_callback = kwargs.get("usage_callback")
|
||||
completion_callback = kwargs.get("completion_callback")
|
||||
@@ -55,9 +55,9 @@ class GeminiClient(BaseAPIClient):
|
||||
messages=messages, # type: ignore
|
||||
stream=True,
|
||||
stream_options={"include_usage": True},
|
||||
temperature=temperature if temperature is not None else NOT_GIVEN,
|
||||
max_tokens=max_tokens if max_tokens is not None else NOT_GIVEN,
|
||||
top_p=kwargs.get("top_p", NOT_GIVEN),
|
||||
temperature=temperature if temperature is not None else omit,
|
||||
max_tokens=max_tokens if max_tokens is not None else omit,
|
||||
top_p=kwargs.get("top_p", omit),
|
||||
)
|
||||
|
||||
final_usage = None
|
||||
|
||||
108
routstr/upstream/count_tokens.py
Normal file
108
routstr/upstream/count_tokens.py
Normal file
@@ -0,0 +1,108 @@
|
||||
"""Local handling of Anthropic ``/v1/messages/count_tokens`` for upstreams
|
||||
that do not natively expose the endpoint.
|
||||
|
||||
Most non-Anthropic upstreams (OpenAI-compat, Gemini OpenAI-compat,
|
||||
OpenRouter chat-completions, generic providers) return 400/404 when asked
|
||||
to ``POST /messages/count_tokens``. Claude Code and other Anthropic SDK
|
||||
clients call this endpoint before each turn to size context windows and
|
||||
trigger compaction, so a failure breaks the whole chat.
|
||||
|
||||
We answer locally. ``litellm.token_counter`` understands the Anthropic
|
||||
message shape and the per-model tokenizers, so we prefer it. If it raises
|
||||
(unknown model, encoding lookup failure, ...), we fall back to the
|
||||
project's own ``estimate_tokens`` heuristic, which is always defined and
|
||||
never raises.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
import litellm
|
||||
from fastapi.responses import Response
|
||||
|
||||
from ..core import get_logger
|
||||
from ..payment.helpers import estimate_tokens
|
||||
from ..payment.models import Model
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
def _parse_request_body(request_body: bytes | None) -> dict[str, Any]:
|
||||
if not request_body:
|
||||
return {}
|
||||
try:
|
||||
parsed = json.loads(request_body)
|
||||
except (ValueError, TypeError):
|
||||
return {}
|
||||
return parsed if isinstance(parsed, dict) else {}
|
||||
|
||||
|
||||
def _count_with_litellm(model: str, body: dict[str, Any]) -> int:
|
||||
messages = body.get("messages")
|
||||
if not isinstance(messages, list):
|
||||
messages = []
|
||||
|
||||
system = body.get("system")
|
||||
if isinstance(system, str) and system:
|
||||
messages = [{"role": "system", "content": system}, *messages]
|
||||
elif isinstance(system, list):
|
||||
text = "".join(
|
||||
block.get("text", "")
|
||||
for block in system
|
||||
if isinstance(block, dict) and block.get("type") == "text"
|
||||
)
|
||||
if text:
|
||||
messages = [{"role": "system", "content": text}, *messages]
|
||||
|
||||
tools = body.get("tools") if isinstance(body.get("tools"), list) else None
|
||||
|
||||
return int(
|
||||
litellm.token_counter(
|
||||
model=model,
|
||||
messages=messages,
|
||||
tools=tools,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def count_tokens_locally(
|
||||
request_body: bytes | None,
|
||||
model_obj: Model | None,
|
||||
) -> Response:
|
||||
"""Return an Anthropic-compatible count_tokens response without
|
||||
touching the upstream. Always returns 200; never raises."""
|
||||
body = _parse_request_body(request_body)
|
||||
|
||||
model_name = ""
|
||||
if model_obj is not None:
|
||||
model_name = model_obj.forwarded_model_id or model_obj.id or ""
|
||||
if not model_name:
|
||||
body_model = body.get("model")
|
||||
if isinstance(body_model, str):
|
||||
model_name = body_model
|
||||
|
||||
input_tokens: int
|
||||
try:
|
||||
input_tokens = _count_with_litellm(model_name, body)
|
||||
except Exception as exc:
|
||||
messages = body.get("messages")
|
||||
fallback_messages = messages if isinstance(messages, list) else []
|
||||
input_tokens = estimate_tokens(fallback_messages)
|
||||
logger.debug(
|
||||
"litellm token_counter failed; using local estimator",
|
||||
extra={
|
||||
"model": model_name,
|
||||
"error": str(exc),
|
||||
"error_type": type(exc).__name__,
|
||||
"estimated_tokens": input_tokens,
|
||||
},
|
||||
)
|
||||
|
||||
payload = {"input_tokens": max(0, int(input_tokens))}
|
||||
return Response(
|
||||
content=json.dumps(payload).encode(),
|
||||
status_code=200,
|
||||
media_type="application/json",
|
||||
)
|
||||
@@ -12,6 +12,7 @@ class FireworksUpstreamProvider(BaseUpstreamProvider):
|
||||
provider_type = "fireworks"
|
||||
default_base_url = "https://api.fireworks.ai/inference/v1"
|
||||
platform_url = "https://app.fireworks.ai/settings/users/api-keys"
|
||||
litellm_provider_prefix = "fireworks_ai/"
|
||||
|
||||
def __init__(self, api_key: str, provider_fee: float = 1.01):
|
||||
super().__init__(
|
||||
|
||||
@@ -1,17 +1,13 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import AsyncGenerator
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from fastapi import Request
|
||||
from fastapi.responses import Response, StreamingResponse
|
||||
|
||||
from . import gemini_messages
|
||||
from .base import BaseUpstreamProvider
|
||||
from .clients.gemini import GeminiClient
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import ApiKey, AsyncSession, UpstreamProviderRow
|
||||
from ..core.db import UpstreamProviderRow
|
||||
from ..payment.models import Model
|
||||
|
||||
from ..core.logging import get_logger
|
||||
@@ -20,9 +16,19 @@ logger = get_logger(__name__)
|
||||
|
||||
|
||||
class GeminiUpstreamProvider(BaseUpstreamProvider):
|
||||
"""Gemini provider — proxies through Gemini's OpenAI-compat surface.
|
||||
|
||||
The chat-completions, embeddings, and models paths all flow through
|
||||
:meth:`BaseUpstreamProvider.forward_request`; we only override
|
||||
``get_request_base_url`` to redirect to ``{base}/openai/...`` and
|
||||
``_dispatch_anthropic_messages`` to inject thought-signatures on the
|
||||
/v1/messages path (see :mod:`gemini_messages` for that rationale).
|
||||
"""
|
||||
|
||||
provider_type = "gemini"
|
||||
default_base_url = "https://generativelanguage.googleapis.com/v1beta"
|
||||
platform_url = "https://aistudio.google.com/app/apikey"
|
||||
litellm_provider_prefix = "gemini/"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -39,7 +45,7 @@ class GeminiUpstreamProvider(BaseUpstreamProvider):
|
||||
|
||||
@property
|
||||
def client(self) -> GeminiClient:
|
||||
"""Get or create the Gemini API client."""
|
||||
"""Get or create the Gemini API client (used for the models listing)."""
|
||||
if self._client is None:
|
||||
self._client = GeminiClient(api_key=self.api_key)
|
||||
return self._client
|
||||
@@ -65,240 +71,66 @@ class GeminiUpstreamProvider(BaseUpstreamProvider):
|
||||
}
|
||||
|
||||
def transform_model_name(self, model_id: str) -> str:
|
||||
return model_id.removeprefix("gemini/")
|
||||
"""Reduce a routstr model id to the bare upstream Gemini name.
|
||||
|
||||
async def forward_request(
|
||||
Gemini's OpenAI-compat surface expects the literal model id
|
||||
(e.g. ``gemini-3.1-flash-lite-preview``) — no ``gemini/`` provider
|
||||
prefix and no ``google/`` vendor sub-prefix. Take the last path
|
||||
segment so we tolerate any of:
|
||||
|
||||
``gemini-2.0-flash``
|
||||
``gemini/gemini-2.0-flash``
|
||||
``gemini/google/gemini-3.1-flash-lite-preview``
|
||||
"""
|
||||
return model_id.rsplit("/", 1)[-1]
|
||||
|
||||
@property
|
||||
def compat_base_url(self) -> str:
|
||||
"""Gemini's OpenAI-compat surface, regardless of what's stored.
|
||||
|
||||
Stored ``base_url`` may be ``.../v1beta`` (the native Gemini API
|
||||
root) or ``.../v1beta/openai`` (already pointed at the compat
|
||||
surface). Normalize to the latter.
|
||||
"""
|
||||
return self.base_url.rstrip("/").removesuffix("/openai") + "/openai"
|
||||
|
||||
def get_request_base_url(
|
||||
self, path: str, model_obj: "Model | None" = None
|
||||
) -> str:
|
||||
"""Route every proxied request to the OpenAI-compat surface.
|
||||
|
||||
Required because the stored ``base_url`` typically points at the
|
||||
native Gemini API (``/v1beta``), but :meth:`forward_request`
|
||||
forwards OpenAI-shaped paths (``/chat/completions``,
|
||||
``/embeddings``, ``/models``) which only exist under the
|
||||
``/openai`` subtree.
|
||||
"""
|
||||
return self.compat_base_url
|
||||
|
||||
async def _dispatch_anthropic_messages(
|
||||
self,
|
||||
request: Request,
|
||||
path: str,
|
||||
headers: dict,
|
||||
request_body: bytes | None,
|
||||
key: ApiKey,
|
||||
max_cost_for_model: int,
|
||||
session: AsyncSession,
|
||||
model_obj: Model,
|
||||
) -> Response | StreamingResponse:
|
||||
# Remove provider prefix from model ID for Gemini API
|
||||
if "/" in model_obj.id:
|
||||
model_obj.id = model_obj.id.split("/", 1)[1]
|
||||
model_obj: "Model",
|
||||
*,
|
||||
log_extra: dict[str, Any] | None = None,
|
||||
) -> tuple[bool, Any, str | None]:
|
||||
"""Dispatch /v1/messages via the gemini-specific httpx path.
|
||||
|
||||
if not path.startswith("chat/completions"):
|
||||
return await super().forward_request(
|
||||
request,
|
||||
path,
|
||||
headers,
|
||||
request_body,
|
||||
key,
|
||||
max_cost_for_model,
|
||||
session,
|
||||
model_obj,
|
||||
)
|
||||
|
||||
if not request_body:
|
||||
return await super().forward_request(
|
||||
request,
|
||||
path,
|
||||
headers,
|
||||
request_body,
|
||||
key,
|
||||
max_cost_for_model,
|
||||
session,
|
||||
model_obj,
|
||||
)
|
||||
|
||||
try:
|
||||
openai_data = json.loads(request_body)
|
||||
messages = openai_data.get("messages", [])
|
||||
temperature = openai_data.get("temperature")
|
||||
max_tokens = openai_data.get("max_tokens")
|
||||
top_p = openai_data.get("top_p")
|
||||
is_streaming = openai_data.get("stream", False)
|
||||
|
||||
logger.info(
|
||||
"Processing Gemini request with client abstraction",
|
||||
extra={
|
||||
"model": model_obj.id,
|
||||
"is_streaming": is_streaming,
|
||||
"message_count": len(messages),
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
},
|
||||
)
|
||||
|
||||
if is_streaming:
|
||||
final_usage_data: dict | None = None
|
||||
|
||||
def usage_callback(usage_data: dict[str, Any]) -> None:
|
||||
"""Callback to capture usage data during streaming"""
|
||||
nonlocal final_usage_data
|
||||
final_usage_data = usage_data
|
||||
|
||||
async def completion_callback(
|
||||
model: str, usage_data: dict[str, Any] | None
|
||||
) -> None:
|
||||
"""Callback to handle payment when streaming completes"""
|
||||
nonlocal final_usage_data
|
||||
if usage_data:
|
||||
final_usage_data = usage_data
|
||||
|
||||
payment_data = {
|
||||
"model": model,
|
||||
"usage": final_usage_data,
|
||||
}
|
||||
|
||||
from ..auth import adjust_payment_for_tokens
|
||||
from ..core.db import create_session
|
||||
|
||||
async with create_session() as new_session:
|
||||
fresh_key = await new_session.get(key.__class__, key.hashed_key)
|
||||
if fresh_key:
|
||||
try:
|
||||
cost_data = await adjust_payment_for_tokens(
|
||||
fresh_key,
|
||||
payment_data,
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"Gemini streaming payment finalized",
|
||||
extra={
|
||||
"cost_data": cost_data,
|
||||
"usage_data": final_usage_data,
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
},
|
||||
)
|
||||
except Exception as cost_error:
|
||||
logger.error(
|
||||
"Error finalizing Gemini streaming payment",
|
||||
extra={
|
||||
"error": str(cost_error),
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
},
|
||||
)
|
||||
|
||||
response_generator = self.client.generate_content_stream(
|
||||
model=model_obj.id,
|
||||
messages=messages,
|
||||
temperature=temperature,
|
||||
max_tokens=max_tokens,
|
||||
top_p=top_p,
|
||||
usage_callback=usage_callback,
|
||||
completion_callback=completion_callback,
|
||||
)
|
||||
|
||||
async def stream_with_cost() -> AsyncGenerator[bytes, None]:
|
||||
payment_finalized = False
|
||||
|
||||
async def finalize_payment() -> None:
|
||||
nonlocal payment_finalized
|
||||
if payment_finalized:
|
||||
return
|
||||
from ..auth import adjust_payment_for_tokens
|
||||
from ..core.db import create_session
|
||||
|
||||
async with create_session() as new_session:
|
||||
fresh_key = await new_session.get(
|
||||
key.__class__, key.hashed_key
|
||||
)
|
||||
if fresh_key:
|
||||
try:
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key,
|
||||
{
|
||||
"model": model_obj.id,
|
||||
"usage": final_usage_data,
|
||||
},
|
||||
new_session,
|
||||
max_cost_for_model,
|
||||
)
|
||||
payment_finalized = True
|
||||
except Exception as cost_error:
|
||||
logger.error(
|
||||
"Error finalizing Gemini streaming payment in fallback",
|
||||
extra={
|
||||
"error": str(cost_error),
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
async for chunk in response_generator:
|
||||
sse_data = f"data: {json.dumps(chunk)}\n\n"
|
||||
yield sse_data.encode()
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Error in Gemini streaming response",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
},
|
||||
)
|
||||
raise
|
||||
finally:
|
||||
if not payment_finalized:
|
||||
await finalize_payment()
|
||||
|
||||
return StreamingResponse(
|
||||
stream_with_cost(),
|
||||
media_type="text/event-stream",
|
||||
headers={"Cache-Control": "no-cache", "Connection": "keep-alive"},
|
||||
)
|
||||
|
||||
else:
|
||||
openai_format_response = await self.client.generate_content(
|
||||
model=model_obj.id,
|
||||
messages=messages,
|
||||
temperature=temperature,
|
||||
max_tokens=max_tokens,
|
||||
top_p=top_p,
|
||||
)
|
||||
|
||||
from ..auth import adjust_payment_for_tokens
|
||||
|
||||
cost_data = await adjust_payment_for_tokens(
|
||||
key, openai_format_response, session, max_cost_for_model
|
||||
)
|
||||
openai_format_response["cost"] = cost_data
|
||||
|
||||
logger.info(
|
||||
"Gemini non-streaming payment completed",
|
||||
extra={
|
||||
"cost_data": cost_data,
|
||||
"model": model_obj.id,
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
},
|
||||
)
|
||||
|
||||
return Response(
|
||||
content=json.dumps(openai_format_response),
|
||||
media_type="application/json",
|
||||
headers={"Cache-Control": "no-cache"},
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Error in Gemini forward_request",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
"path": path,
|
||||
"key_hash": key.hashed_key[:8] + "...",
|
||||
},
|
||||
)
|
||||
return await super().forward_request(
|
||||
request,
|
||||
path,
|
||||
headers,
|
||||
request_body,
|
||||
key,
|
||||
max_cost_for_model,
|
||||
session,
|
||||
model_obj,
|
||||
)
|
||||
See :mod:`routstr.upstream.gemini_messages` for the full rationale
|
||||
(thought-signature injection, why litellm + the openai SDK can't
|
||||
carry the required ``extra_content`` field).
|
||||
"""
|
||||
return await gemini_messages.dispatch_gemini_messages(
|
||||
request_body=request_body,
|
||||
model_obj=model_obj,
|
||||
base_url=self.compat_base_url,
|
||||
api_key=self.api_key,
|
||||
transform_model_name=self.transform_model_name,
|
||||
log_extra=log_extra,
|
||||
)
|
||||
|
||||
async def _fetch_provider_models(self) -> dict:
|
||||
"""Fetch models from Gemini API."""
|
||||
"""Fetch models from Gemini API via the OpenAI-compat client."""
|
||||
try:
|
||||
models_data = await self.client.list_models()
|
||||
|
||||
|
||||
480
routstr/upstream/gemini_messages.py
Normal file
480
routstr/upstream/gemini_messages.py
Normal file
@@ -0,0 +1,480 @@
|
||||
"""Custom /v1/messages dispatcher for Gemini's OpenAI-compat endpoint.
|
||||
|
||||
Why this exists
|
||||
---------------
|
||||
|
||||
Gemini 2.5 / 3 thinking models reject inbound ``functionCall`` parts that
|
||||
lack a ``thought_signature`` field once any prior turn in the conversation
|
||||
contains a function call. Anthropic-Messages clients (Claude Code etc.)
|
||||
have no concept of thought signatures, so multi-turn tool conversations
|
||||
fail with::
|
||||
|
||||
Function call is missing a thought_signature in functionCall parts.
|
||||
|
||||
Google's published escape hatch (https://ai.google.dev/gemini-api/docs/
|
||||
thought-signatures, FAQ #1) is the dummy signature
|
||||
``"skip_thought_signature_validator"`` placed at
|
||||
``tool_calls[i].extra_content.google.thought_signature`` for every tool
|
||||
call in the request. The hatch is documented specifically for
|
||||
"transferring a trace from a different model that does not include thought
|
||||
signatures" — exactly our case.
|
||||
|
||||
Why we can't reach the wire via litellm
|
||||
---------------------------------------
|
||||
|
||||
``litellm.anthropic.messages.acreate`` flows through the openai SDK, whose
|
||||
pydantic ``ChatCompletionMessageToolCall`` model silently drops unknown
|
||||
fields like ``extra_content``. Litellm has no openai-compat translator
|
||||
that emits ``extra_content.google.thought_signature``. So we bypass both
|
||||
litellm and the openai SDK at the transport layer.
|
||||
|
||||
Pipeline
|
||||
--------
|
||||
|
||||
1. Translate Anthropic body → OpenAI body via litellm's
|
||||
``AnthropicAdapter`` (the same translator
|
||||
``litellm.anthropic.messages.acreate`` uses internally).
|
||||
2. Inject ``extra_content.google.thought_signature`` on every
|
||||
``tool_calls[]`` entry.
|
||||
3. Set ``reasoning_effort="none"`` to disable Gemini's thinking pass.
|
||||
4. POST directly to ``{base_url}/chat/completions`` with ``stream=true``
|
||||
via ``httpx`` (preserves arbitrary fields verbatim).
|
||||
5. Translate OpenAI streaming chunks → Anthropic SSE events.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from collections.abc import AsyncGenerator, AsyncIterator
|
||||
from typing import Any, Callable
|
||||
|
||||
import httpx
|
||||
|
||||
from ..core import get_logger
|
||||
from ..core.exceptions import UpstreamError
|
||||
from ..payment.models import Model
|
||||
from .messages_dispatch import (
|
||||
ANTHROPIC_ONLY_FIELDS,
|
||||
aggregate_anthropic_events_to_message,
|
||||
)
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
DUMMY_THOUGHT_SIGNATURE = "skip_thought_signature_validator"
|
||||
|
||||
# Mapping: OpenAI finish_reason → Anthropic stop_reason
|
||||
_FINISH_TO_STOP = {
|
||||
"stop": "end_turn",
|
||||
"length": "max_tokens",
|
||||
"tool_calls": "tool_use",
|
||||
"function_call": "tool_use",
|
||||
"content_filter": "refusal",
|
||||
}
|
||||
|
||||
|
||||
def inject_thought_signatures(messages: list[dict]) -> None:
|
||||
"""Add ``extra_content.google.thought_signature`` to every tool_call.
|
||||
|
||||
Mutates ``messages`` in place. Idempotent: existing signatures are not
|
||||
overwritten.
|
||||
"""
|
||||
for msg in messages:
|
||||
tool_calls = msg.get("tool_calls")
|
||||
if not isinstance(tool_calls, list):
|
||||
continue
|
||||
for tc in tool_calls:
|
||||
if not isinstance(tc, dict):
|
||||
continue
|
||||
extra = tc.get("extra_content")
|
||||
if not isinstance(extra, dict):
|
||||
extra = tc["extra_content"] = {}
|
||||
google_cfg = extra.get("google")
|
||||
if not isinstance(google_cfg, dict):
|
||||
google_cfg = extra["google"] = {}
|
||||
google_cfg.setdefault("thought_signature", DUMMY_THOUGHT_SIGNATURE)
|
||||
|
||||
|
||||
def _translate_anthropic_to_openai(body: dict, model: str) -> dict:
|
||||
"""Use litellm's translator to convert an Anthropic /messages body to
|
||||
OpenAI /chat/completions kwargs.
|
||||
|
||||
Imported lazily because the litellm internal path is heavy and not
|
||||
needed for any other code path in routstr.
|
||||
"""
|
||||
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import ( # noqa: E501
|
||||
AnthropicAdapter,
|
||||
)
|
||||
|
||||
kwargs = {"model": model, **body}
|
||||
translated = AnthropicAdapter().translate_completion_input_params(kwargs)
|
||||
if translated is None:
|
||||
raise UpstreamError(
|
||||
"Failed to translate Anthropic body to OpenAI format",
|
||||
status_code=500,
|
||||
)
|
||||
return dict(translated)
|
||||
|
||||
|
||||
def _sse_event(event_type: str, payload: dict) -> bytes:
|
||||
return f"event: {event_type}\ndata: {json.dumps(payload)}\n\n".encode()
|
||||
|
||||
|
||||
async def _openai_chunks_to_anthropic_events(
|
||||
line_iter: AsyncIterator[str], requested_model: str | None
|
||||
) -> AsyncGenerator[bytes, None]:
|
||||
"""Translate an OpenAI chat-completions SSE byte stream into the
|
||||
Anthropic-Messages SSE event sequence.
|
||||
|
||||
Maintains per-chunk state across:
|
||||
|
||||
* one optional text content block (lazy-opened on first text delta)
|
||||
* any number of tool_use blocks indexed by openai's ``delta.tool_calls[].index``
|
||||
* final ``stop_reason`` / ``usage`` carried out via ``message_delta`` /
|
||||
``message_stop``
|
||||
"""
|
||||
msg_id = f"msg_{uuid.uuid4().hex[:24]}"
|
||||
started = False
|
||||
text_block_idx: int | None = None
|
||||
tool_block_indices: dict[int, int] = {}
|
||||
next_block_idx = 0
|
||||
final_finish_reason: str | None = None
|
||||
final_usage: dict[str, int] = {"input_tokens": 0, "output_tokens": 0}
|
||||
|
||||
def open_text_block() -> bytes:
|
||||
nonlocal text_block_idx, next_block_idx
|
||||
text_block_idx = next_block_idx
|
||||
next_block_idx += 1
|
||||
return _sse_event(
|
||||
"content_block_start",
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": text_block_idx,
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
},
|
||||
)
|
||||
|
||||
def open_tool_block(delta_idx: int, tc: dict) -> bytes:
|
||||
nonlocal next_block_idx
|
||||
block_idx = next_block_idx
|
||||
next_block_idx += 1
|
||||
tool_block_indices[delta_idx] = block_idx
|
||||
fn = tc.get("function") or {}
|
||||
return _sse_event(
|
||||
"content_block_start",
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": block_idx,
|
||||
"content_block": {
|
||||
"type": "tool_use",
|
||||
"id": tc.get("id") or f"toolu_{uuid.uuid4().hex[:24]}",
|
||||
"name": fn.get("name") or "",
|
||||
"input": {},
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
def close_block(idx: int) -> bytes:
|
||||
return _sse_event(
|
||||
"content_block_stop",
|
||||
{"type": "content_block_stop", "index": idx},
|
||||
)
|
||||
|
||||
async for raw_line in line_iter:
|
||||
line = raw_line.strip()
|
||||
if not line:
|
||||
continue
|
||||
if not line.startswith("data:"):
|
||||
continue
|
||||
payload = line[5:].lstrip()
|
||||
if not payload or payload == "[DONE]":
|
||||
continue
|
||||
try:
|
||||
chunk = json.loads(payload)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if not isinstance(chunk, dict):
|
||||
continue
|
||||
|
||||
if not started:
|
||||
started = True
|
||||
yield _sse_event(
|
||||
"message_start",
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": chunk.get("id") or msg_id,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": requested_model or chunk.get("model") or "",
|
||||
"content": [],
|
||||
"stop_reason": None,
|
||||
"stop_sequence": None,
|
||||
"usage": {
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0,
|
||||
},
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
usage = chunk.get("usage")
|
||||
if isinstance(usage, dict):
|
||||
in_tok = usage.get("prompt_tokens") or usage.get("input_tokens") or 0
|
||||
out_tok = usage.get("completion_tokens") or usage.get("output_tokens") or 0
|
||||
if in_tok:
|
||||
final_usage["input_tokens"] = int(in_tok)
|
||||
if out_tok:
|
||||
final_usage["output_tokens"] = int(out_tok)
|
||||
|
||||
choices = chunk.get("choices") or []
|
||||
if not choices:
|
||||
continue
|
||||
choice = choices[0] if isinstance(choices[0], dict) else {}
|
||||
delta = choice.get("delta") or {}
|
||||
if not isinstance(delta, dict):
|
||||
delta = {}
|
||||
|
||||
# Text delta
|
||||
text = delta.get("content")
|
||||
if isinstance(text, str) and text:
|
||||
if text_block_idx is None:
|
||||
yield open_text_block()
|
||||
yield _sse_event(
|
||||
"content_block_delta",
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": text_block_idx,
|
||||
"delta": {"type": "text_delta", "text": text},
|
||||
},
|
||||
)
|
||||
|
||||
# Tool call deltas
|
||||
tool_calls_delta = delta.get("tool_calls")
|
||||
if isinstance(tool_calls_delta, list):
|
||||
for tc in tool_calls_delta:
|
||||
if not isinstance(tc, dict):
|
||||
continue
|
||||
d_idx = int(tc.get("index") or 0)
|
||||
if d_idx not in tool_block_indices:
|
||||
yield open_tool_block(d_idx, tc)
|
||||
block_idx = tool_block_indices[d_idx]
|
||||
fn = tc.get("function") or {}
|
||||
args = fn.get("arguments")
|
||||
if isinstance(args, str) and args:
|
||||
yield _sse_event(
|
||||
"content_block_delta",
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": block_idx,
|
||||
"delta": {
|
||||
"type": "input_json_delta",
|
||||
"partial_json": args,
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
finish = choice.get("finish_reason")
|
||||
if finish:
|
||||
final_finish_reason = finish
|
||||
|
||||
# Close any open content blocks
|
||||
if text_block_idx is not None:
|
||||
yield close_block(text_block_idx)
|
||||
for block_idx in tool_block_indices.values():
|
||||
yield close_block(block_idx)
|
||||
|
||||
# message_delta with stop_reason and usage
|
||||
stop_reason = _FINISH_TO_STOP.get(final_finish_reason or "", "end_turn")
|
||||
yield _sse_event(
|
||||
"message_delta",
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": stop_reason, "stop_sequence": None},
|
||||
"usage": final_usage,
|
||||
},
|
||||
)
|
||||
yield _sse_event("message_stop", {"type": "message_stop"})
|
||||
|
||||
|
||||
async def _post_and_stream(
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
payload: dict,
|
||||
log_extra: dict[str, Any] | None,
|
||||
) -> tuple[httpx.AsyncClient, httpx.Response]:
|
||||
"""POST to upstream chat-completions and return (client, response) for
|
||||
streaming. Caller is responsible for closing both."""
|
||||
url = f"{base_url.rstrip('/')}/chat/completions"
|
||||
client = httpx.AsyncClient(timeout=httpx.Timeout(120.0, read=120.0))
|
||||
try:
|
||||
request = client.build_request(
|
||||
"POST",
|
||||
url,
|
||||
json=payload,
|
||||
headers={
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "text/event-stream",
|
||||
},
|
||||
)
|
||||
response = await client.send(request, stream=True)
|
||||
except Exception as exc:
|
||||
await client.aclose()
|
||||
logger.error(
|
||||
"Gemini messages dispatch HTTP error",
|
||||
extra={"error": str(exc), "url": url, **(log_extra or {})},
|
||||
)
|
||||
raise UpstreamError(
|
||||
f"Failed to reach Gemini upstream: {exc}", status_code=502
|
||||
) from exc
|
||||
|
||||
if response.status_code >= 400:
|
||||
try:
|
||||
body_bytes = await response.aread()
|
||||
finally:
|
||||
await response.aclose()
|
||||
await client.aclose()
|
||||
body_text = body_bytes.decode("utf-8", errors="replace")
|
||||
logger.error(
|
||||
"Gemini messages dispatch upstream error",
|
||||
extra={
|
||||
"status_code": response.status_code,
|
||||
"body": body_text[:1000],
|
||||
"url": url,
|
||||
**(log_extra or {}),
|
||||
},
|
||||
)
|
||||
raise UpstreamError(
|
||||
f"Upstream error via gemini compat: {body_text}",
|
||||
status_code=response.status_code,
|
||||
)
|
||||
|
||||
return client, response
|
||||
|
||||
|
||||
async def dispatch_gemini_messages(
|
||||
*,
|
||||
request_body: bytes | None,
|
||||
model_obj: Model,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
transform_model_name: Callable[[str], str],
|
||||
log_extra: dict[str, Any] | None = None,
|
||||
) -> tuple[bool, Any, str | None]:
|
||||
"""Dispatch a /v1/messages request to Gemini's OpenAI-compat endpoint
|
||||
with thought-signature injection.
|
||||
|
||||
Returns ``(client_stream, result, requested_model)`` where ``result``
|
||||
is either an ``AsyncIterator[bytes]`` of Anthropic-format SSE events
|
||||
(for streaming clients) or an Anthropic Message dict (after the caller
|
||||
aggregates).
|
||||
"""
|
||||
if not request_body:
|
||||
raise UpstreamError(
|
||||
"Missing request body for /v1/messages", status_code=400
|
||||
)
|
||||
|
||||
try:
|
||||
body: dict = json.loads(request_body)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise UpstreamError(
|
||||
f"Invalid JSON in /v1/messages body: {exc}", status_code=400
|
||||
) from exc
|
||||
|
||||
body.pop("model", None)
|
||||
client_stream = bool(body.pop("stream", False))
|
||||
|
||||
# Anthropic-Messages-only fields that don't translate to OpenAI
|
||||
# Chat Completions. litellm's translator passes through unknown
|
||||
# top-level fields verbatim and Gemini's compat surface 400s on
|
||||
# unknown names like ``context_management`` / ``output_config``.
|
||||
dropped: dict[str, Any] = {}
|
||||
for field in ANTHROPIC_ONLY_FIELDS:
|
||||
if field in body:
|
||||
dropped[field] = body.pop(field)
|
||||
if dropped:
|
||||
logger.debug(
|
||||
"Dropped anthropic-only fields before gemini compat dispatch",
|
||||
extra={"dropped_keys": sorted(dropped.keys())},
|
||||
)
|
||||
|
||||
requested_model = (
|
||||
(model_obj.forwarded_model_id or model_obj.id) if model_obj else None
|
||||
)
|
||||
upstream_model = transform_model_name(model_obj.id)
|
||||
|
||||
openai_kwargs = _translate_anthropic_to_openai(body, upstream_model)
|
||||
|
||||
messages = openai_kwargs.get("messages") or []
|
||||
if isinstance(messages, list):
|
||||
inject_thought_signatures(messages)
|
||||
|
||||
# Disable Gemini's thinking pass; the dummy signature already lifts
|
||||
# validation, but skipping thinking entirely avoids degraded model
|
||||
# output and keeps tool-calling deterministic.
|
||||
openai_kwargs.setdefault("reasoning_effort", "none")
|
||||
openai_kwargs["stream"] = True
|
||||
openai_kwargs["model"] = upstream_model
|
||||
# OpenAI-compat backends (including Gemini's) only emit a final
|
||||
# ``usage`` chunk when the request opts in via this flag. Without it
|
||||
# the cost-calculation pipeline can't read real token counts and
|
||||
# falls back to MaxCostData billing.
|
||||
existing_stream_options = openai_kwargs.get("stream_options")
|
||||
merged_stream_options = (
|
||||
dict(existing_stream_options)
|
||||
if isinstance(existing_stream_options, dict)
|
||||
else {}
|
||||
)
|
||||
merged_stream_options.setdefault("include_usage", True)
|
||||
openai_kwargs["stream_options"] = merged_stream_options
|
||||
|
||||
logger.info(
|
||||
"Dispatching /v1/messages via gemini compat (httpx)",
|
||||
extra={
|
||||
"model": upstream_model,
|
||||
"client_stream": client_stream,
|
||||
"messages_with_tool_calls": sum(
|
||||
1 for m in messages if isinstance(m, dict) and m.get("tool_calls")
|
||||
),
|
||||
**(log_extra or {}),
|
||||
},
|
||||
)
|
||||
|
||||
http_client, response = await _post_and_stream(
|
||||
base_url, api_key, openai_kwargs, log_extra
|
||||
)
|
||||
|
||||
async def line_iter() -> AsyncGenerator[str, None]:
|
||||
try:
|
||||
async for line in response.aiter_lines():
|
||||
yield line
|
||||
finally:
|
||||
await response.aclose()
|
||||
await http_client.aclose()
|
||||
|
||||
anthropic_event_iter = _openai_chunks_to_anthropic_events(
|
||||
line_iter(), requested_model
|
||||
)
|
||||
|
||||
if not client_stream:
|
||||
# Aggregate the Anthropic SSE byte stream into a single Message dict
|
||||
# so the rest of the pipeline (cost calc, metadata injection,
|
||||
# response building) can treat it identically to a non-streaming
|
||||
# litellm response.
|
||||
try:
|
||||
aggregated = await aggregate_anthropic_events_to_message(
|
||||
anthropic_event_iter
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.error(
|
||||
"Failed to aggregate Gemini compat events into message",
|
||||
extra={"error": str(exc), **(log_extra or {})},
|
||||
)
|
||||
raise UpstreamError(
|
||||
f"Failed to aggregate upstream stream: {exc}",
|
||||
status_code=502,
|
||||
) from exc
|
||||
return client_stream, aggregated, requested_model
|
||||
|
||||
return client_stream, anthropic_event_iter, requested_model
|
||||
@@ -12,6 +12,7 @@ class GroqUpstreamProvider(BaseUpstreamProvider):
|
||||
provider_type = "groq"
|
||||
default_base_url = "https://api.groq.com/openai/v1"
|
||||
platform_url = "https://console.groq.com/keys"
|
||||
litellm_provider_prefix = "groq/"
|
||||
|
||||
def __init__(self, api_key: str, provider_fee: float = 1.01):
|
||||
super().__init__(
|
||||
|
||||
@@ -3,7 +3,7 @@ from __future__ import annotations
|
||||
import asyncio
|
||||
import os
|
||||
import re
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Callable
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.settings import Settings
|
||||
@@ -122,12 +122,16 @@ async def get_all_models_with_overrides(
|
||||
|
||||
|
||||
async def refresh_upstreams_models_periodically(
|
||||
upstreams: list[BaseUpstreamProvider],
|
||||
upstreams_provider: (
|
||||
Callable[[], list[BaseUpstreamProvider]] | list[BaseUpstreamProvider]
|
||||
),
|
||||
) -> None:
|
||||
"""Background task to periodically refresh models cache for all providers.
|
||||
|
||||
Args:
|
||||
upstreams: List of upstream provider instances
|
||||
upstreams_provider: Either a callable returning the live upstream list
|
||||
(preferred — picks up providers added/changed via reinitialize_upstreams),
|
||||
or a static list (legacy, will go stale after reinitialize_upstreams).
|
||||
"""
|
||||
import asyncio
|
||||
import random
|
||||
@@ -139,9 +143,14 @@ async def refresh_upstreams_models_periodically(
|
||||
logger.info("Provider models refresh disabled (interval <= 0)")
|
||||
return
|
||||
|
||||
def _resolve_upstreams() -> list[BaseUpstreamProvider]:
|
||||
if callable(upstreams_provider):
|
||||
return upstreams_provider()
|
||||
return upstreams_provider
|
||||
|
||||
while True:
|
||||
try:
|
||||
for upstream in upstreams:
|
||||
for upstream in _resolve_upstreams():
|
||||
try:
|
||||
await upstream.refresh_models_cache()
|
||||
except Exception as e:
|
||||
@@ -188,13 +197,14 @@ async def init_upstreams() -> list[BaseUpstreamProvider]:
|
||||
existing_providers = result.all()
|
||||
|
||||
if not existing_providers:
|
||||
logger.info(
|
||||
"No upstream providers found in database, seeding from settings"
|
||||
)
|
||||
await _seed_providers_from_settings(session, settings)
|
||||
await session.commit()
|
||||
result = await session.exec(select(UpstreamProviderRow))
|
||||
existing_providers = result.all()
|
||||
if existing_providers:
|
||||
logger.info(
|
||||
f"Seeded {len(existing_providers)} upstream providers from settings"
|
||||
)
|
||||
|
||||
async def _init_single_provider(
|
||||
provider_row: UpstreamProviderRow,
|
||||
@@ -205,6 +215,9 @@ async def init_upstreams() -> list[BaseUpstreamProvider]:
|
||||
|
||||
provider = _instantiate_provider(provider_row)
|
||||
if provider:
|
||||
# Keep provider DB id on runtime instance so model mapping can
|
||||
# bind DB overrides to the correct upstream.
|
||||
setattr(provider, "db_id", provider_row.id)
|
||||
await provider.refresh_models_cache()
|
||||
logger.debug(
|
||||
f"Initialized {provider_row.provider_type} provider",
|
||||
@@ -236,7 +249,7 @@ async def _seed_providers_from_settings(
|
||||
from . import upstream_provider_classes
|
||||
|
||||
providers_to_add: list[UpstreamProviderRow] = []
|
||||
seeded_base_urls: set[str] = set()
|
||||
seeded_provider_keys: set[tuple[str, str]] = set()
|
||||
|
||||
provider_classes_by_type = {
|
||||
cls.provider_type: cls
|
||||
@@ -261,7 +274,8 @@ async def _seed_providers_from_settings(
|
||||
base_url = provider_class.default_base_url # type: ignore[attr-defined]
|
||||
result = await session.exec(
|
||||
select(UpstreamProviderRow).where(
|
||||
UpstreamProviderRow.base_url == base_url
|
||||
UpstreamProviderRow.base_url == base_url,
|
||||
UpstreamProviderRow.api_key == api_key,
|
||||
)
|
||||
)
|
||||
if not result.first():
|
||||
@@ -273,13 +287,15 @@ async def _seed_providers_from_settings(
|
||||
enabled=True,
|
||||
)
|
||||
)
|
||||
seeded_base_urls.add(base_url)
|
||||
seeded_provider_keys.add((base_url, api_key))
|
||||
|
||||
ollama_base_url = os.environ.get("OLLAMA_BASE_URL")
|
||||
if ollama_base_url:
|
||||
ollama_api_key = os.environ.get("OLLAMA_API_KEY", "")
|
||||
result = await session.exec(
|
||||
select(UpstreamProviderRow).where(
|
||||
UpstreamProviderRow.base_url == ollama_base_url
|
||||
UpstreamProviderRow.base_url == ollama_base_url,
|
||||
UpstreamProviderRow.api_key == ollama_api_key,
|
||||
)
|
||||
)
|
||||
if not result.first():
|
||||
@@ -287,18 +303,20 @@ async def _seed_providers_from_settings(
|
||||
UpstreamProviderRow(
|
||||
provider_type="ollama",
|
||||
base_url=ollama_base_url,
|
||||
api_key=os.environ.get("OLLAMA_API_KEY", ""),
|
||||
api_key=ollama_api_key,
|
||||
enabled=True,
|
||||
)
|
||||
)
|
||||
seeded_base_urls.add(ollama_base_url)
|
||||
seeded_provider_keys.add((ollama_base_url, ollama_api_key))
|
||||
|
||||
if settings.chat_completions_api_version and settings.upstream_base_url:
|
||||
base_url = settings.upstream_base_url
|
||||
if base_url not in seeded_base_urls:
|
||||
api_key = settings.upstream_api_key
|
||||
if (base_url, api_key) not in seeded_provider_keys:
|
||||
result = await session.exec(
|
||||
select(UpstreamProviderRow).where(
|
||||
UpstreamProviderRow.base_url == base_url
|
||||
UpstreamProviderRow.base_url == base_url,
|
||||
UpstreamProviderRow.api_key == api_key,
|
||||
)
|
||||
)
|
||||
if not result.first():
|
||||
@@ -306,19 +324,21 @@ async def _seed_providers_from_settings(
|
||||
UpstreamProviderRow(
|
||||
provider_type="azure",
|
||||
base_url=base_url,
|
||||
api_key=settings.upstream_api_key,
|
||||
api_key=api_key,
|
||||
api_version=settings.chat_completions_api_version,
|
||||
enabled=True,
|
||||
)
|
||||
)
|
||||
seeded_base_urls.add(base_url)
|
||||
seeded_provider_keys.add((base_url, api_key))
|
||||
|
||||
if settings.upstream_base_url and settings.upstream_api_key:
|
||||
base_url = settings.upstream_base_url
|
||||
if base_url not in seeded_base_urls:
|
||||
api_key = settings.upstream_api_key
|
||||
if (base_url, api_key) not in seeded_provider_keys:
|
||||
result = await session.exec(
|
||||
select(UpstreamProviderRow).where(
|
||||
UpstreamProviderRow.base_url == base_url
|
||||
UpstreamProviderRow.base_url == base_url,
|
||||
UpstreamProviderRow.api_key == api_key,
|
||||
)
|
||||
)
|
||||
if not result.first():
|
||||
@@ -326,11 +346,11 @@ async def _seed_providers_from_settings(
|
||||
UpstreamProviderRow(
|
||||
provider_type="custom",
|
||||
base_url=base_url,
|
||||
api_key=settings.upstream_api_key,
|
||||
api_key=api_key,
|
||||
enabled=True,
|
||||
)
|
||||
)
|
||||
seeded_base_urls.add(base_url)
|
||||
seeded_provider_keys.add((base_url, api_key))
|
||||
|
||||
for provider in providers_to_add:
|
||||
session.add(provider)
|
||||
|
||||
162
routstr/upstream/litellm_routing.py
Normal file
162
routstr/upstream/litellm_routing.py
Normal file
@@ -0,0 +1,162 @@
|
||||
"""Map an upstream `base_url` to the correct litellm provider prefix.
|
||||
|
||||
Used by `BaseUpstreamProvider.get_litellm_provider_prefix` so that custom /
|
||||
generic provider rows (which inherit the base class default) get routed to
|
||||
the right litellm backend instead of falling back to `openai/`.
|
||||
|
||||
The table is compiled from litellm 1.74's
|
||||
`litellm/litellm_core_utils/get_llm_provider_logic.py` (the
|
||||
`openai_compatible_endpoints` table) plus the providers documented at
|
||||
https://docs.litellm.ai/docs/providers. Substring match is used so that
|
||||
URLs with paths, ports, regional subdomains, etc. all resolve correctly.
|
||||
|
||||
Order matters: more specific needles must appear before more generic ones
|
||||
(e.g. ``openai.azure.com`` before ``api.openai.com``).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import litellm
|
||||
|
||||
DEFAULT_PREFIX = "openai/"
|
||||
|
||||
LITELLM_HOST_PREFIX_MAP: tuple[tuple[str, str], ...] = (
|
||||
# Azure must win over api.openai.com because the host ends with
|
||||
# `openai.azure.com` and we don't want it picked up as plain OpenAI.
|
||||
("openai.azure.com", "azure/"),
|
||||
# Google
|
||||
("generativelanguage.googleapis.com", "gemini/"),
|
||||
("aiplatform.googleapis.com", "vertex_ai/"),
|
||||
# First-class providers with native litellm prefixes
|
||||
("api.openai.com", "openai/"),
|
||||
("api.anthropic.com", "anthropic/"),
|
||||
("api.groq.com", "groq/"),
|
||||
("api.fireworks.ai", "fireworks_ai/"),
|
||||
("api.x.ai", "xai/"),
|
||||
("api.perplexity.ai", "perplexity/"),
|
||||
("openrouter.ai", "openrouter/"),
|
||||
("api.deepseek.com", "deepseek/"),
|
||||
("api.together.xyz", "together_ai/"),
|
||||
("codestral.mistral.ai", "codestral/"),
|
||||
("api.mistral.ai", "mistral/"),
|
||||
("api.cohere.com", "cohere_chat/"),
|
||||
("api.cohere.ai", "cohere_chat/"),
|
||||
("api.deepinfra.com", "deepinfra/"),
|
||||
("api.endpoints.anyscale.com", "anyscale/"),
|
||||
("api.cerebras.ai", "cerebras/"),
|
||||
("inference.baseten.co", "baseten/"),
|
||||
("api.sambanova.ai", "sambanova/"),
|
||||
("api.ai21.com", "ai21_chat/"),
|
||||
("api.friendli.ai", "friendliai/"),
|
||||
("api.galadriel.com", "galadriel/"),
|
||||
("api.llama.com", "meta_llama/"),
|
||||
("api.featherless.ai", "featherless_ai/"),
|
||||
("inference.api.nscale.com", "nscale/"),
|
||||
("dashscope-intl.aliyuncs.com", "dashscope/"),
|
||||
("api.moonshot.ai", "moonshot/"),
|
||||
("api.moonshot.cn", "moonshot/"),
|
||||
("api.minimax.io", "minimax/"),
|
||||
("api.minimaxi.com", "minimax/"),
|
||||
("platform.publicai.co", "publicai/"),
|
||||
("api.synthetic.new", "synthetic/"),
|
||||
("api.stima.tech", "apertis/"),
|
||||
("nano-gpt.com", "nano-gpt/"),
|
||||
("api.poe.com", "poe/"),
|
||||
("llm.chutes.ai", "chutes/"),
|
||||
("api.v0.dev", "v0/"),
|
||||
("api.lambda.ai", "lambda_ai/"),
|
||||
("api.hyperbolic.xyz", "hyperbolic/"),
|
||||
("ai-gateway.vercel.sh", "vercel_ai_gateway/"),
|
||||
("api.inference.wandb.ai", "wandb/"),
|
||||
("integrate.api.nvidia.com", "nvidia_nim/"),
|
||||
("api.studio.nebius.com", "nebius/"),
|
||||
("api.novita.ai", "novita/"),
|
||||
("ark.cn-beijing.volces.com", "volcengine/"),
|
||||
("api.voyageai.com", "voyage/"),
|
||||
("api.jina.ai", "jina_ai/"),
|
||||
("api.aimlapi.com", "aiml/"),
|
||||
("api.snowflakecomputing.com", "snowflake/"),
|
||||
("databricks.com", "databricks/"),
|
||||
("huggingface.co", "huggingface/"),
|
||||
)
|
||||
|
||||
# Substrings that indicate an Ollama deployment regardless of port/scheme.
|
||||
OLLAMA_HOST_HINTS: tuple[str, ...] = (
|
||||
"localhost:11434",
|
||||
"127.0.0.1:11434",
|
||||
"ollama",
|
||||
)
|
||||
|
||||
|
||||
def detect_litellm_prefix(
|
||||
base_url: str | None, default: str = DEFAULT_PREFIX
|
||||
) -> str:
|
||||
"""Return the litellm provider prefix (`"<provider>/"`) for `base_url`.
|
||||
|
||||
Falls back to `default` when the host doesn't match any known provider.
|
||||
The default is `openai/` because every unmatched OpenAI-compatible
|
||||
server is, by definition, an OpenAI-compatible server.
|
||||
"""
|
||||
if not base_url:
|
||||
return default
|
||||
|
||||
parsed = urlsplit(base_url)
|
||||
host = parsed.netloc.lower() or base_url.lower()
|
||||
|
||||
for needle, prefix in LITELLM_HOST_PREFIX_MAP:
|
||||
if needle in host:
|
||||
return prefix
|
||||
|
||||
if any(hint in host for hint in OLLAMA_HOST_HINTS):
|
||||
return "ollama_chat/"
|
||||
|
||||
return default
|
||||
|
||||
|
||||
_configured = False
|
||||
|
||||
|
||||
def configure_litellm() -> None:
|
||||
"""Apply litellm global settings used by the messages-dispatch path.
|
||||
|
||||
Idempotent: safe to call from both app startup and module-level
|
||||
initializers without side effects on the second invocation.
|
||||
|
||||
Settings applied:
|
||||
|
||||
* ``LITELLM_DEBUG=1`` enables litellm's verbose debug logger.
|
||||
* Forces the Anthropic-messages adapter to call OpenAI Chat Completions
|
||||
(POST ``/chat/completions``) instead of the Responses API (POST
|
||||
``/responses``) for ``openai/``-prefixed providers. OpenAI-compatible
|
||||
upstreams like Google's generativelanguage compat endpoint expose
|
||||
``/chat/completions`` but not ``/responses``, which would 404. Set
|
||||
``LITELLM_USE_RESPONSES_API_FOR_ANTHROPIC_MESSAGES=1`` to opt out.
|
||||
* Silently drops Anthropic-Messages-only parameters (``thinking``,
|
||||
``cache_control``, ``context_management``, ...) when translating to
|
||||
providers that don't accept them, instead of raising
|
||||
``UnsupportedParamsError``. Set ``LITELLM_STRICT_PARAMS=1`` to opt
|
||||
out.
|
||||
"""
|
||||
global _configured
|
||||
if _configured:
|
||||
return
|
||||
|
||||
if os.getenv("LITELLM_DEBUG") == "1":
|
||||
try:
|
||||
litellm._turn_on_debug() # type: ignore[no-untyped-call]
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if os.getenv("LITELLM_USE_RESPONSES_API_FOR_ANTHROPIC_MESSAGES") != "1":
|
||||
try:
|
||||
litellm.use_chat_completions_url_for_anthropic_messages = True
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if os.getenv("LITELLM_STRICT_PARAMS") != "1":
|
||||
litellm.drop_params = True
|
||||
|
||||
_configured = True
|
||||
561
routstr/upstream/messages_dispatch.py
Normal file
561
routstr/upstream/messages_dispatch.py
Normal file
@@ -0,0 +1,561 @@
|
||||
"""Pure helpers for translating ``/v1/messages`` to upstream chat completions
|
||||
via litellm.
|
||||
|
||||
This module owns the litellm/Anthropic-Messages translation layer:
|
||||
|
||||
* SSE parsing (``parse_sse_blocks``, ``events_from_chunk``)
|
||||
* Payload coercion (``coerce_litellm_payload``)
|
||||
* Stream aggregation (``aggregate_anthropic_events_to_message``) — drains
|
||||
a streamed Anthropic event sequence into a single Message dict
|
||||
* Per-event annotation for streaming (``annotate_event``,
|
||||
``stream_annotated_events``) — handles the model-rewrite + token-tally
|
||||
bookkeeping shared by the bearer-key and x-cashu streaming paths
|
||||
* The dispatch entry point (``dispatch_anthropic_messages``)
|
||||
* Refund math (``compute_refund``)
|
||||
|
||||
Nothing in here touches ``BaseUpstreamProvider``; the thin instance methods
|
||||
on the provider class forward to these functions and only retain logic that
|
||||
genuinely needs ``self`` (cost adjustment, metadata injection, refund
|
||||
sending).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import AsyncGenerator, AsyncIterator
|
||||
from typing import Any, Callable, NamedTuple, cast
|
||||
|
||||
import litellm
|
||||
|
||||
from ..core import get_logger
|
||||
from ..core.exceptions import UpstreamError
|
||||
from ..payment.models import Model
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
# Anthropic-Messages-only fields that don't translate to OpenAI
|
||||
# Chat Completions. ``litellm.drop_params`` only filters *known*
|
||||
# unsupported params; these newer/extension fields get passed through
|
||||
# verbatim and the upstream rejects them with a 400. Pop them here so the
|
||||
# request reaches the upstream cleanly.
|
||||
ANTHROPIC_ONLY_FIELDS: tuple[str, ...] = (
|
||||
"thinking",
|
||||
"cache_control",
|
||||
"context_management",
|
||||
"output_config",
|
||||
"mcp_servers",
|
||||
"service_tier",
|
||||
"anthropic_version",
|
||||
"anthropic_beta",
|
||||
)
|
||||
|
||||
|
||||
def coerce_litellm_payload(payload: object) -> dict:
|
||||
"""Convert a litellm event into a plain dict.
|
||||
|
||||
Non-streaming responses come back as Anthropic-shaped pydantic models
|
||||
or dicts. Streaming may yield raw bytes/str (SSE-encoded); those go
|
||||
through ``events_from_chunk`` instead, not here.
|
||||
"""
|
||||
if isinstance(payload, dict):
|
||||
return dict(payload)
|
||||
if hasattr(payload, "model_dump"):
|
||||
return cast(dict, payload.model_dump())
|
||||
raise TypeError(f"Cannot coerce {type(payload).__name__} to dict")
|
||||
|
||||
|
||||
def parse_sse_blocks(buffer: bytes) -> tuple[list[dict], bytes]:
|
||||
"""Parse complete SSE event blocks out of a byte buffer.
|
||||
|
||||
Returns (events, remaining_buffer). Events are JSON objects parsed from
|
||||
one or more ``data:`` lines per block. Comments, blank lines, and
|
||||
``[DONE]`` sentinels are ignored. A trailing partial block is preserved
|
||||
in remaining_buffer.
|
||||
"""
|
||||
events: list[dict] = []
|
||||
while True:
|
||||
sep = buffer.find(b"\n\n")
|
||||
if sep < 0:
|
||||
sep_rn = buffer.find(b"\r\n\r\n")
|
||||
if sep_rn < 0:
|
||||
break
|
||||
block = buffer[:sep_rn]
|
||||
buffer = buffer[sep_rn + 4 :]
|
||||
else:
|
||||
block = buffer[:sep]
|
||||
buffer = buffer[sep + 2 :]
|
||||
|
||||
data_lines: list[str] = []
|
||||
for raw_line in block.replace(b"\r\n", b"\n").split(b"\n"):
|
||||
line = raw_line.decode("utf-8", errors="replace")
|
||||
if line.startswith(":"):
|
||||
continue
|
||||
if line.startswith("data:"):
|
||||
data_lines.append(line[5:].lstrip())
|
||||
if not data_lines:
|
||||
continue
|
||||
payload = "\n".join(data_lines).strip()
|
||||
if not payload or payload == "[DONE]":
|
||||
continue
|
||||
try:
|
||||
obj = json.loads(payload)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if isinstance(obj, dict):
|
||||
events.append(obj)
|
||||
return events, buffer
|
||||
|
||||
|
||||
def events_from_chunk(
|
||||
chunk: object, sse_buffer: bytes
|
||||
) -> tuple[list[dict], bytes]:
|
||||
"""Normalize a stream chunk into one or more event dicts.
|
||||
|
||||
``litellm.anthropic.messages.acreate(stream=True)`` yields raw SSE
|
||||
bytes in practice; some adapters yield strings or typed events. Handle
|
||||
all three.
|
||||
"""
|
||||
if isinstance(chunk, (bytes, bytearray)):
|
||||
sse_buffer += bytes(chunk)
|
||||
events, sse_buffer = parse_sse_blocks(sse_buffer)
|
||||
return events, sse_buffer
|
||||
if isinstance(chunk, str):
|
||||
sse_buffer += chunk.encode("utf-8")
|
||||
events, sse_buffer = parse_sse_blocks(sse_buffer)
|
||||
return events, sse_buffer
|
||||
return [coerce_litellm_payload(chunk)], sse_buffer
|
||||
|
||||
|
||||
async def aggregate_anthropic_events_to_message(
|
||||
iterator: AsyncIterator[Any],
|
||||
) -> dict:
|
||||
"""Drain an Anthropic-Messages event iterator into a single Message dict.
|
||||
|
||||
Produces the shape ``litellm.anthropic.messages.acreate(stream=False)``
|
||||
would have returned. Used to transparently stream from upstream while
|
||||
still returning a non-streaming response to the client. Lets us
|
||||
sidestep upstream quirks (e.g. Fireworks rejects ``max_tokens > 4096``
|
||||
unless ``stream=true``) without leaking that into client-visible
|
||||
behavior.
|
||||
"""
|
||||
sse_buffer = b""
|
||||
message: dict = {}
|
||||
blocks: list[dict] = []
|
||||
partial_json: dict[int, str] = {}
|
||||
final_stop_reason: str | None = None
|
||||
final_stop_sequence: str | None = None
|
||||
final_usage: dict[str, Any] = {}
|
||||
final_model: str | None = None
|
||||
|
||||
async for chunk in iterator:
|
||||
events, sse_buffer = events_from_chunk(chunk, sse_buffer)
|
||||
for event in events:
|
||||
etype = event.get("type")
|
||||
if etype == "message_start":
|
||||
raw = event.get("message") or {}
|
||||
if isinstance(raw, dict):
|
||||
message = dict(raw)
|
||||
existing = message.get("content")
|
||||
blocks = list(existing) if isinstance(existing, list) else []
|
||||
usage = message.get("usage")
|
||||
if isinstance(usage, dict):
|
||||
final_usage = dict(usage)
|
||||
if isinstance(message.get("model"), str):
|
||||
final_model = message["model"]
|
||||
elif etype == "content_block_start":
|
||||
idx = int(event.get("index") or 0)
|
||||
cb = event.get("content_block") or {}
|
||||
cb_dict = dict(cb) if isinstance(cb, dict) else {}
|
||||
while len(blocks) <= idx:
|
||||
blocks.append({})
|
||||
blocks[idx] = cb_dict
|
||||
elif etype == "content_block_delta":
|
||||
idx = int(event.get("index") or 0)
|
||||
if idx >= len(blocks):
|
||||
continue
|
||||
delta = event.get("delta") or {}
|
||||
if not isinstance(delta, dict):
|
||||
continue
|
||||
dtype = delta.get("type")
|
||||
block = blocks[idx]
|
||||
if dtype == "text_delta":
|
||||
block["text"] = (block.get("text") or "") + (
|
||||
delta.get("text") or ""
|
||||
)
|
||||
elif dtype == "input_json_delta":
|
||||
partial_json[idx] = partial_json.get(idx, "") + (
|
||||
delta.get("partial_json") or ""
|
||||
)
|
||||
elif dtype == "thinking_delta":
|
||||
block["thinking"] = (block.get("thinking") or "") + (
|
||||
delta.get("thinking") or ""
|
||||
)
|
||||
elif dtype == "signature_delta":
|
||||
block["signature"] = (block.get("signature") or "") + (
|
||||
delta.get("signature") or ""
|
||||
)
|
||||
elif etype == "content_block_stop":
|
||||
idx = int(event.get("index") or 0)
|
||||
raw_json = partial_json.pop(idx, None)
|
||||
if raw_json is not None and idx < len(blocks):
|
||||
try:
|
||||
blocks[idx]["input"] = (
|
||||
json.loads(raw_json) if raw_json else {}
|
||||
)
|
||||
except json.JSONDecodeError:
|
||||
blocks[idx]["input"] = raw_json
|
||||
elif etype == "message_delta":
|
||||
delta = event.get("delta") or {}
|
||||
if isinstance(delta, dict):
|
||||
if "stop_reason" in delta:
|
||||
final_stop_reason = delta.get("stop_reason")
|
||||
if "stop_sequence" in delta:
|
||||
final_stop_sequence = delta.get("stop_sequence")
|
||||
usage = event.get("usage")
|
||||
if isinstance(usage, dict):
|
||||
final_usage.update(usage)
|
||||
# message_stop: nothing to merge
|
||||
|
||||
if not message:
|
||||
# Upstream returned no message_start; expose what we can so the
|
||||
# client at least sees the assembled content.
|
||||
message = {
|
||||
"id": "",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
}
|
||||
|
||||
message["content"] = blocks
|
||||
if final_model and not message.get("model"):
|
||||
message["model"] = final_model
|
||||
if final_stop_reason is not None:
|
||||
message["stop_reason"] = final_stop_reason
|
||||
if final_stop_sequence is not None:
|
||||
message["stop_sequence"] = final_stop_sequence
|
||||
if final_usage:
|
||||
existing_usage = message.get("usage")
|
||||
merged = dict(existing_usage) if isinstance(existing_usage, dict) else {}
|
||||
merged.update(final_usage)
|
||||
message["usage"] = merged
|
||||
return message
|
||||
|
||||
|
||||
class AnnotatedEvent(NamedTuple):
|
||||
"""One Anthropic SSE event after model-rewrite + token-tally bookkeeping.
|
||||
|
||||
``sse_bytes`` is the wire-ready ``event:`` / ``data:`` block; the two
|
||||
streaming paths in ``BaseUpstreamProvider`` consume ``sse_bytes`` plus
|
||||
the tallies and only differ in whether they stream live or buffer
|
||||
first.
|
||||
|
||||
``cache_read_input_tokens`` and ``cache_creation_input_tokens`` are
|
||||
surfaced separately so the cost path can price them against the cache
|
||||
rate rather than fold them silently into the regular input bucket.
|
||||
|
||||
``total_cost`` / ``input_cost`` / ``output_cost`` carry any
|
||||
USD cost figures the upstream attached to this event (from
|
||||
``usage.cost``, ``usage.total_cost``, or ``usage.cost_details``) so the
|
||||
streaming paths can re-embed them in the rebuilt ``usage`` dict and let
|
||||
``calculate_cost`` convert directly USD→sats instead of falling back to
|
||||
token-based math.
|
||||
"""
|
||||
|
||||
event: dict
|
||||
sse_bytes: bytes
|
||||
input_tokens: int
|
||||
output_tokens: int
|
||||
cache_read_input_tokens: int
|
||||
cache_creation_input_tokens: int
|
||||
total_cost: float
|
||||
input_cost: float
|
||||
output_cost: float
|
||||
model: str | None
|
||||
|
||||
|
||||
def annotate_event(event: dict, requested_model: str | None) -> AnnotatedEvent:
|
||||
"""Rewrite ``model`` fields and extract per-event token / model info.
|
||||
|
||||
Mutates ``event`` in place when ``requested_model`` is set so the
|
||||
upstream's true model name doesn't leak to the client.
|
||||
"""
|
||||
if requested_model:
|
||||
msg = event.get("message")
|
||||
if isinstance(msg, dict) and "model" in msg:
|
||||
msg["model"] = requested_model
|
||||
if "model" in event:
|
||||
event["model"] = requested_model
|
||||
|
||||
in_tokens = 0
|
||||
out_tokens = 0
|
||||
cache_read_tokens = 0
|
||||
cache_create_tokens = 0
|
||||
total_cost = 0.0
|
||||
input_cost = 0.0
|
||||
output_cost = 0.0
|
||||
model: str | None = None
|
||||
|
||||
def _coerce_float(value: object) -> float:
|
||||
if value is None or isinstance(value, bool):
|
||||
return 0.0
|
||||
if not isinstance(value, (int, float, str)):
|
||||
return 0.0
|
||||
try:
|
||||
return max(0.0, float(value))
|
||||
except (TypeError, ValueError):
|
||||
return 0.0
|
||||
|
||||
def _accumulate(usage: dict) -> None:
|
||||
nonlocal in_tokens, out_tokens, cache_read_tokens, cache_create_tokens
|
||||
nonlocal total_cost, input_cost, output_cost
|
||||
in_tokens += int(usage.get("input_tokens") or 0)
|
||||
out_tokens += int(usage.get("output_tokens") or 0)
|
||||
cache_read_tokens += int(usage.get("cache_read_input_tokens") or 0)
|
||||
cache_create_tokens += int(usage.get("cache_creation_input_tokens") or 0)
|
||||
total_cost += _coerce_float(usage.get("total_cost"))
|
||||
input_cost += _coerce_float(usage.get("input_cost"))
|
||||
output_cost += _coerce_float(usage.get("output_cost"))
|
||||
|
||||
msg_for_meta = event.get("message")
|
||||
if isinstance(msg_for_meta, dict):
|
||||
if msg_for_meta.get("model"):
|
||||
model = str(msg_for_meta["model"])
|
||||
usage = msg_for_meta.get("usage")
|
||||
if isinstance(usage, dict):
|
||||
_accumulate(usage)
|
||||
|
||||
if isinstance(event.get("usage"), dict):
|
||||
_accumulate(event["usage"])
|
||||
|
||||
# Some upstreams (notably OpenRouter-style proxies) attach cost fields
|
||||
# directly at the event root rather than inside ``usage``.
|
||||
for field in ("total_cost", "cost"):
|
||||
total_cost = max(total_cost, _coerce_float(event.get(field)))
|
||||
input_cost = max(input_cost, _coerce_float(event.get("input_cost")))
|
||||
output_cost = max(output_cost, _coerce_float(event.get("output_cost")))
|
||||
root_cost_details = event.get("cost_details")
|
||||
if isinstance(root_cost_details, dict):
|
||||
total_cost = max(
|
||||
total_cost,
|
||||
_coerce_float(root_cost_details.get("total_cost")),
|
||||
)
|
||||
input_cost = max(
|
||||
input_cost,
|
||||
_coerce_float(root_cost_details.get("input_cost")),
|
||||
)
|
||||
output_cost = max(
|
||||
output_cost,
|
||||
_coerce_float(root_cost_details.get("output_cost")),
|
||||
)
|
||||
|
||||
event_type = str(event.get("type") or "")
|
||||
payload = json.dumps(event)
|
||||
if event_type:
|
||||
sse_bytes = f"event: {event_type}\ndata: {payload}\n\n".encode()
|
||||
else:
|
||||
sse_bytes = f"data: {payload}\n\n".encode()
|
||||
|
||||
return AnnotatedEvent(
|
||||
event,
|
||||
sse_bytes,
|
||||
in_tokens,
|
||||
out_tokens,
|
||||
cache_read_tokens,
|
||||
cache_create_tokens,
|
||||
total_cost,
|
||||
input_cost,
|
||||
output_cost,
|
||||
model,
|
||||
)
|
||||
|
||||
|
||||
async def stream_annotated_events(
|
||||
iterator: AsyncIterator[Any],
|
||||
requested_model: str | None,
|
||||
) -> AsyncGenerator[AnnotatedEvent, None]:
|
||||
"""Yield annotated, SSE-serialized events from a litellm stream.
|
||||
|
||||
Both streaming paths in ``BaseUpstreamProvider`` consume this; the only
|
||||
divergence between them — yield-as-you-go vs buffer-then-replay — stays
|
||||
in the caller.
|
||||
"""
|
||||
sse_buffer = b""
|
||||
async for chunk in iterator:
|
||||
events, sse_buffer = events_from_chunk(chunk, sse_buffer)
|
||||
for event in events:
|
||||
yield annotate_event(event, requested_model)
|
||||
|
||||
|
||||
def embed_usd_costs(
|
||||
usage: dict,
|
||||
total_cost: float,
|
||||
input_cost: float,
|
||||
output_cost: float,
|
||||
) -> None:
|
||||
"""Mutate ``usage`` so ``calculate_cost`` will pick up the USD totals.
|
||||
|
||||
Mirrors the upstream shape: when any USD figure is present, attach
|
||||
``cost`` (used by the simple-fallback branch in ``calculate_cost``) and
|
||||
a ``cost_details`` block (used by the preferred branch — also gives the
|
||||
input/output USD split when we have one).
|
||||
"""
|
||||
if total_cost <= 0 and input_cost <= 0 and output_cost <= 0:
|
||||
return
|
||||
|
||||
cost_details: dict[str, float] = {}
|
||||
effective_total = total_cost
|
||||
if effective_total <= 0 and (input_cost > 0 or output_cost > 0):
|
||||
effective_total = input_cost + output_cost
|
||||
|
||||
if effective_total > 0:
|
||||
cost_details["total_cost"] = effective_total
|
||||
usage["cost"] = effective_total
|
||||
if input_cost > 0:
|
||||
cost_details["input_cost"] = input_cost
|
||||
if output_cost > 0:
|
||||
cost_details["output_cost"] = output_cost
|
||||
if cost_details:
|
||||
usage["cost_details"] = cost_details
|
||||
|
||||
|
||||
def compute_refund(amount: int, unit: str, cost_msats: int) -> int:
|
||||
if unit == "msat":
|
||||
return amount - cost_msats
|
||||
if unit == "sat":
|
||||
return amount - (cost_msats + 999) // 1000
|
||||
raise ValueError(f"Invalid unit: {unit}")
|
||||
|
||||
|
||||
async def dispatch_anthropic_messages(
|
||||
*,
|
||||
request_body: bytes | None,
|
||||
model_obj: Model,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
provider_prefix: str,
|
||||
transform_model_name: Callable[[str], str],
|
||||
log_extra: dict[str, Any] | None = None,
|
||||
) -> tuple[bool, Any, str | None]:
|
||||
"""Call ``litellm.anthropic.messages.acreate`` and return
|
||||
``(client_stream, result, requested_model)``.
|
||||
|
||||
Shared by the bearer-key and x-cashu paths. Raises :class:`UpstreamError`
|
||||
on bad input or upstream failure.
|
||||
"""
|
||||
if not request_body:
|
||||
raise UpstreamError(
|
||||
"Missing request body for /v1/messages", status_code=400
|
||||
)
|
||||
|
||||
try:
|
||||
body: dict = json.loads(request_body)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise UpstreamError(
|
||||
f"Invalid JSON in /v1/messages body: {exc}", status_code=400
|
||||
) from exc
|
||||
|
||||
body.pop("model", None)
|
||||
# `stream` here is what the **client** asked for. Upstream is always
|
||||
# streamed (see `upstream_stream` below); when the client asked for a
|
||||
# non-streaming response we drain and aggregate the events into a
|
||||
# single Anthropic Message dict before returning. This sidesteps
|
||||
# provider-specific non-streaming caps (e.g. Fireworks rejects
|
||||
# `max_tokens > 4096` unless `stream=true`).
|
||||
client_stream = bool(body.pop("stream", False))
|
||||
upstream_stream = True
|
||||
|
||||
dropped: dict[str, Any] = {}
|
||||
for field in ANTHROPIC_ONLY_FIELDS:
|
||||
if field in body:
|
||||
dropped[field] = body.pop(field)
|
||||
if dropped:
|
||||
logger.debug(
|
||||
"Dropped anthropic-only fields before litellm dispatch",
|
||||
extra={"dropped_keys": sorted(dropped.keys())},
|
||||
)
|
||||
|
||||
# Convention: `model.id` is the canonical upstream model name;
|
||||
# `forwarded_model_id` is the public alias the internal API exposes
|
||||
# and echoes back to the client.
|
||||
requested_model = (
|
||||
(model_obj.forwarded_model_id or model_obj.id) if model_obj else None
|
||||
)
|
||||
upstream_model = transform_model_name(model_obj.id)
|
||||
litellm_model = f"{provider_prefix}{upstream_model}"
|
||||
|
||||
kwargs: dict = {
|
||||
"model": litellm_model,
|
||||
"api_base": base_url,
|
||||
"api_key": api_key,
|
||||
"stream": upstream_stream,
|
||||
**body,
|
||||
}
|
||||
|
||||
logger.info(
|
||||
"Dispatching /v1/messages via litellm",
|
||||
extra={
|
||||
"model": litellm_model,
|
||||
"resolved_provider": provider_prefix.rstrip("/"),
|
||||
"client_stream": client_stream,
|
||||
"upstream_stream": upstream_stream,
|
||||
**(log_extra or {}),
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
result = await litellm.anthropic.messages.acreate(**kwargs)
|
||||
except Exception as exc:
|
||||
exc_message = getattr(exc, "message", None) or str(exc) or repr(exc)
|
||||
exc_status = getattr(exc, "status_code", None)
|
||||
exc_response = getattr(exc, "response", None)
|
||||
response_text = None
|
||||
if exc_response is not None:
|
||||
try:
|
||||
response_text = getattr(exc_response, "text", str(exc_response))
|
||||
except Exception:
|
||||
response_text = "<unreadable>"
|
||||
logger.error(
|
||||
"litellm dispatch failed",
|
||||
extra={
|
||||
"error": exc_message,
|
||||
"error_type": type(exc).__name__,
|
||||
"status_code": exc_status,
|
||||
"llm_provider": getattr(exc, "llm_provider", None),
|
||||
"body": getattr(exc, "body", None),
|
||||
"response_text": response_text,
|
||||
"model": litellm_model,
|
||||
"api_base": base_url,
|
||||
},
|
||||
)
|
||||
raise UpstreamError(
|
||||
f"Upstream error via litellm: {exc_message}",
|
||||
status_code=exc_status if isinstance(exc_status, int) else 502,
|
||||
) from exc
|
||||
|
||||
if not client_stream and hasattr(result, "__aiter__"):
|
||||
# Client asked for a non-streaming response but we always stream
|
||||
# from upstream — drain the events into a single Anthropic Message
|
||||
# dict so the rest of the pipeline can treat it as if upstream had
|
||||
# returned non-streaming. Some litellm adapters return a
|
||||
# non-streaming dict even when ``stream=True``; in that case,
|
||||
# leave the result as-is.
|
||||
try:
|
||||
aggregated: Any = await aggregate_anthropic_events_to_message(
|
||||
cast(AsyncIterator[Any], result)
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.error(
|
||||
"Failed to aggregate streamed events into message",
|
||||
extra={
|
||||
"error": str(exc),
|
||||
"error_type": type(exc).__name__,
|
||||
"model": litellm_model,
|
||||
},
|
||||
)
|
||||
raise UpstreamError(
|
||||
f"Failed to aggregate upstream stream: {exc}",
|
||||
status_code=502,
|
||||
) from exc
|
||||
return client_stream, aggregated, requested_model
|
||||
|
||||
return client_stream, result, requested_model
|
||||
@@ -3,13 +3,11 @@ from __future__ import annotations
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import httpx
|
||||
from fastapi import Request
|
||||
from fastapi.responses import Response, StreamingResponse
|
||||
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import ApiKey, AsyncSession, UpstreamProviderRow
|
||||
from ..core.db import UpstreamProviderRow
|
||||
from ..payment.models import Model
|
||||
|
||||
from ..core.logging import get_logger
|
||||
@@ -23,6 +21,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
|
||||
provider_type = "ollama"
|
||||
default_base_url = "http://localhost:11434"
|
||||
platform_url = None
|
||||
litellm_provider_prefix = "ollama_chat/"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -67,38 +66,11 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
|
||||
"""Strip 'ollama/' prefix for Ollama API compatibility."""
|
||||
return model_id.removeprefix("ollama/")
|
||||
|
||||
async def forward_request(
|
||||
self,
|
||||
request: Request,
|
||||
path: str,
|
||||
headers: dict,
|
||||
request_body: bytes | None,
|
||||
key: ApiKey,
|
||||
max_cost_for_model: int,
|
||||
session: AsyncSession,
|
||||
model_obj: Model,
|
||||
) -> Response | StreamingResponse:
|
||||
"""Override to use OpenAI-compatible endpoint for proxy requests."""
|
||||
if path.startswith("v1/"):
|
||||
path = path.replace("v1/", "")
|
||||
|
||||
original_base_url = self.base_url
|
||||
self.base_url = f"{self.base_url}/v1"
|
||||
|
||||
try:
|
||||
result = await super().forward_request(
|
||||
request,
|
||||
path,
|
||||
headers,
|
||||
request_body,
|
||||
key,
|
||||
max_cost_for_model,
|
||||
session,
|
||||
model_obj,
|
||||
)
|
||||
return result
|
||||
finally:
|
||||
self.base_url = original_base_url
|
||||
def get_request_base_url(
|
||||
self, path: str, model_obj: Model | None = None
|
||||
) -> str:
|
||||
"""Route proxy traffic through Ollama's OpenAI-compatible /v1 endpoint."""
|
||||
return f"{self.base_url.rstrip('/')}/v1"
|
||||
|
||||
async def fetch_models(self) -> list[Model]:
|
||||
"""Fetch models from Ollama API using /api/tags endpoint."""
|
||||
@@ -213,7 +185,7 @@ class OllamaUpstreamProvider(BaseUpstreamProvider):
|
||||
except Exception:
|
||||
self._models_cache = models_with_fees
|
||||
|
||||
self._models_by_id = {m.id: m for m in self._models_cache}
|
||||
self._models_by_id = {m.forwarded_model_id or m.id: m for m in self._models_cache}
|
||||
logger.info(
|
||||
f"Refreshed models cache for {self.base_url}",
|
||||
extra={"model_count": len(models)},
|
||||
|
||||
@@ -15,6 +15,8 @@ class OpenRouterUpstreamProvider(BaseUpstreamProvider):
|
||||
provider_type = "openrouter"
|
||||
default_base_url = "https://openrouter.ai/api/v1"
|
||||
platform_url = "https://openrouter.ai/settings/keys"
|
||||
supports_anthropic_messages = True
|
||||
litellm_provider_prefix = "openrouter/"
|
||||
|
||||
def __init__(self, api_key: str, provider_fee: float = 1.06):
|
||||
"""Initialize OpenRouter provider with API key.
|
||||
|
||||
@@ -13,6 +13,7 @@ class PerplexityUpstreamProvider(BaseUpstreamProvider):
|
||||
provider_type = "perplexity"
|
||||
default_base_url = "https://api.perplexity.ai/"
|
||||
platform_url = "https://www.perplexity.ai/account/api/keys"
|
||||
litellm_provider_prefix = "perplexity/"
|
||||
|
||||
def __init__(self, api_key: str, provider_fee: float = 1.01):
|
||||
super().__init__(
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel
|
||||
from pydantic.v1 import BaseModel, Field
|
||||
|
||||
from ..core.logging import get_logger
|
||||
from ..payment.models import Architecture, Model, Pricing, async_fetch_openrouter_models
|
||||
@@ -16,18 +16,20 @@ logger = get_logger(__name__)
|
||||
|
||||
|
||||
class PPQAIModelPricing(BaseModel):
|
||||
ui: dict[str, float]
|
||||
api: dict[str, float]
|
||||
ui: Optional[dict[str, float]] = None
|
||||
api: Optional[dict[str, float]] = None
|
||||
input_per_1M_tokens: Optional[float] = Field(None, alias="input_per_1M_tokens")
|
||||
output_per_1M_tokens: Optional[float] = Field(None, alias="output_per_1M_tokens")
|
||||
|
||||
|
||||
class PPQAIModel(BaseModel):
|
||||
id: str
|
||||
provider: str
|
||||
provider: Optional[str] = None
|
||||
name: str
|
||||
created_at: int
|
||||
context_length: int
|
||||
pricing: PPQAIModelPricing
|
||||
popular: bool
|
||||
popular: bool = False
|
||||
|
||||
|
||||
class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
@@ -134,31 +136,54 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
)
|
||||
|
||||
if or_model:
|
||||
if input_price := ppqai_model.pricing.api.get(
|
||||
"input_per_1M"
|
||||
):
|
||||
input_price = None
|
||||
if ppqai_model.pricing.api:
|
||||
input_price = ppqai_model.pricing.api.get(
|
||||
"input_per_1M"
|
||||
)
|
||||
elif ppqai_model.pricing.input_per_1M_tokens:
|
||||
input_price = ppqai_model.pricing.input_per_1M_tokens
|
||||
|
||||
if input_price is not None:
|
||||
or_model.pricing.prompt = input_price / 1_000_000
|
||||
if output_price := ppqai_model.pricing.api.get(
|
||||
"output_per_1M"
|
||||
):
|
||||
|
||||
output_price = None
|
||||
if ppqai_model.pricing.api:
|
||||
output_price = ppqai_model.pricing.api.get(
|
||||
"output_per_1M"
|
||||
)
|
||||
elif ppqai_model.pricing.output_per_1M_tokens:
|
||||
output_price = ppqai_model.pricing.output_per_1M_tokens
|
||||
|
||||
if output_price is not None:
|
||||
or_model.pricing.completion = output_price / 1_000_000
|
||||
|
||||
if cl := ppqai_model.context_length:
|
||||
or_model.context_length = cl
|
||||
models.append(or_model)
|
||||
else:
|
||||
input_price = ppqai_model.pricing.api.get(
|
||||
"input_per_1M", 0.0
|
||||
)
|
||||
output_price = ppqai_model.pricing.api.get(
|
||||
"output_per_1M", 0.0
|
||||
)
|
||||
input_price = 0.0
|
||||
if ppqai_model.pricing.api:
|
||||
input_price = ppqai_model.pricing.api.get(
|
||||
"input_per_1M", 0.0
|
||||
)
|
||||
elif ppqai_model.pricing.input_per_1M_tokens:
|
||||
input_price = ppqai_model.pricing.input_per_1M_tokens
|
||||
|
||||
output_price = 0.0
|
||||
if ppqai_model.pricing.api:
|
||||
output_price = ppqai_model.pricing.api.get(
|
||||
"output_per_1M", 0.0
|
||||
)
|
||||
elif ppqai_model.pricing.output_per_1M_tokens:
|
||||
output_price = ppqai_model.pricing.output_per_1M_tokens
|
||||
|
||||
models.append(
|
||||
Model(
|
||||
id=ppqai_model.id,
|
||||
name=ppqai_model.name,
|
||||
created=ppqai_model.created_at // 1000,
|
||||
description=f"{ppqai_model.provider} model",
|
||||
description=f"{ppqai_model.provider or 'PPQ.AI'} model",
|
||||
context_length=ppqai_model.context_length,
|
||||
architecture=Architecture(
|
||||
modality="text->text",
|
||||
@@ -196,6 +221,37 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
)
|
||||
return []
|
||||
|
||||
async def on_upstream_error_redirect(
|
||||
self, status_code: int, error_message: str
|
||||
) -> None:
|
||||
if "insufficient balance" in error_message.lower():
|
||||
logger.warning(
|
||||
f"Disabling PPQ.AI provider ({self.base_url}) due to insufficient balance",
|
||||
extra={"error": error_message},
|
||||
)
|
||||
from sqlmodel import select
|
||||
|
||||
from ..core.db import UpstreamProviderRow, create_session
|
||||
|
||||
async with create_session() as session:
|
||||
statement = select(UpstreamProviderRow).where(
|
||||
UpstreamProviderRow.base_url == self.base_url,
|
||||
UpstreamProviderRow.api_key == self.api_key,
|
||||
)
|
||||
result = await session.exec(statement)
|
||||
provider = result.first()
|
||||
|
||||
if provider:
|
||||
provider.enabled = False
|
||||
session.add(provider)
|
||||
await session.commit()
|
||||
|
||||
# Trigger re-initialization of providers
|
||||
# Import here to avoid circular dependency
|
||||
from ..proxy import reinitialize_upstreams
|
||||
|
||||
await reinitialize_upstreams()
|
||||
|
||||
async def create_account(self) -> dict[str, object]:
|
||||
"""Create a new PPQ.AI account.
|
||||
|
||||
@@ -252,7 +308,6 @@ class PPQAIUpstreamProvider(BaseUpstreamProvider):
|
||||
)
|
||||
|
||||
async with httpx.AsyncClient(timeout=30.0) as client:
|
||||
print(f"Payload: {payload}", "sending to", url)
|
||||
response = await client.post(url, headers=headers, json=payload)
|
||||
response.raise_for_status()
|
||||
invoice_data = response.json()
|
||||
|
||||
180
routstr/upstream/routstr.py
Normal file
180
routstr/upstream/routstr.py
Normal file
@@ -0,0 +1,180 @@
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import httpx
|
||||
|
||||
from ..core import get_logger
|
||||
from ..payment.models import Model
|
||||
from .base import BaseUpstreamProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..core.db import UpstreamProviderRow
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class RoutstrUpstreamProvider(BaseUpstreamProvider):
|
||||
"""Upstream provider for communicating with another Routstr instance."""
|
||||
|
||||
provider_type = "routstr"
|
||||
default_base_url = None
|
||||
platform_url = None
|
||||
# Upstream Routstr nodes serve `/v1/messages` natively, so forward the
|
||||
# request as-is instead of round-tripping through litellm's
|
||||
# Anthropic→OpenAI translator.
|
||||
supports_anthropic_messages = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
provider_fee: float = 1.01,
|
||||
provider_settings: dict | None = None,
|
||||
):
|
||||
"""Initialize Routstr provider.
|
||||
|
||||
Args:
|
||||
base_url: Base URL of the upstream Routstr instance
|
||||
api_key: API key for the upstream Routstr instance
|
||||
provider_fee: Provider fee multiplier
|
||||
provider_settings: Provider-specific settings (auto-topup, etc.)
|
||||
"""
|
||||
# Ensure base_url doesn't end with /v1 as BaseUpstreamProvider appends it if needed
|
||||
# but Routstr paths are usually absolute from base.
|
||||
super().__init__(
|
||||
base_url=base_url.rstrip("/"),
|
||||
api_key=api_key,
|
||||
provider_fee=provider_fee,
|
||||
)
|
||||
self.settings = provider_settings or {}
|
||||
|
||||
def normalize_request_path(
|
||||
self, path: str, model_obj: "Model | None" = None
|
||||
) -> str:
|
||||
"""Preserve the ``v1/`` prefix when forwarding to an upstream Routstr.
|
||||
"""
|
||||
return path.lstrip("/")
|
||||
|
||||
@classmethod
|
||||
def from_db_row(
|
||||
cls, provider_row: "UpstreamProviderRow"
|
||||
) -> "RoutstrUpstreamProvider":
|
||||
import json
|
||||
|
||||
settings = {}
|
||||
if provider_row.provider_settings:
|
||||
try:
|
||||
settings = json.loads(provider_row.provider_settings)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return cls(
|
||||
base_url=provider_row.base_url,
|
||||
api_key=provider_row.api_key,
|
||||
provider_fee=provider_row.provider_fee,
|
||||
provider_settings=settings,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_provider_metadata(cls) -> dict[str, object]:
|
||||
return {
|
||||
"id": cls.provider_type,
|
||||
"name": "Routstr Node",
|
||||
"default_base_url": "",
|
||||
"fixed_base_url": False,
|
||||
"platform_url": cls.platform_url,
|
||||
"can_create_account": False,
|
||||
"can_topup": True,
|
||||
"can_show_balance": True,
|
||||
}
|
||||
|
||||
async def get_balance(self) -> float | None:
|
||||
"""Fetch balance from the upstream Routstr node.
|
||||
|
||||
Returns:
|
||||
Balance in satoshis, or None if failed
|
||||
"""
|
||||
url = f"{self.base_url}/v1/balance/info"
|
||||
headers = {"Authorization": f"Bearer {self.api_key}"} if self.api_key else {}
|
||||
|
||||
async with httpx.AsyncClient() as client:
|
||||
try:
|
||||
response = await client.get(url, headers=headers, timeout=10.0)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
# Routstr balance info usually contains 'balance' in msats or sats
|
||||
# Check for msats and convert to sats
|
||||
if "balance_msats" in data:
|
||||
return float(data["balance_msats"]) / 1000.0
|
||||
return float(data.get("balance", 0))
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to fetch balance from upstream Routstr",
|
||||
extra={"url": url, "error": str(e)},
|
||||
)
|
||||
return None
|
||||
|
||||
async def topup(self, cashu_token: str) -> dict[str, Any]:
|
||||
"""Top up balance on the upstream Routstr node.
|
||||
|
||||
Args:
|
||||
cashu_token: Cashu token to deposit
|
||||
|
||||
Returns:
|
||||
Dict containing top-up result
|
||||
"""
|
||||
url = f"{self.base_url}/v1/balance/topup"
|
||||
headers = {"Authorization": f"Bearer {self.api_key}"}
|
||||
payload = {"cashu_token": cashu_token}
|
||||
|
||||
async with httpx.AsyncClient() as client:
|
||||
try:
|
||||
response = await client.post(
|
||||
url, headers=headers, json=payload, timeout=30.0
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to topup upstream Routstr",
|
||||
extra={"url": url, "error": str(e)},
|
||||
)
|
||||
return {"error": str(e)}
|
||||
|
||||
async def fetch_models(self) -> list[Model]:
|
||||
"""Fetch models from the upstream Routstr node."""
|
||||
url = f"{self.base_url}/v1/models"
|
||||
|
||||
async with httpx.AsyncClient() as client:
|
||||
try:
|
||||
response = await client.get(url, headers={}, timeout=15.0)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
models = data.get("data", [])
|
||||
return [Model(**m) for m in models]
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to fetch models from upstream Routstr",
|
||||
extra={"url": url, "error": str(e)},
|
||||
)
|
||||
return []
|
||||
|
||||
async def refund_balance(self) -> dict[str, Any]:
|
||||
"""Request a refund from the upstream Routstr node.
|
||||
|
||||
Returns:
|
||||
Dict containing refund result and token
|
||||
"""
|
||||
url = f"{self.base_url}/v1/balance/refund"
|
||||
headers = {"Authorization": f"Bearer {self.api_key}"}
|
||||
|
||||
async with httpx.AsyncClient() as client:
|
||||
try:
|
||||
response = await client.post(url, headers=headers, timeout=30.0)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to request refund from upstream Routstr",
|
||||
extra={"url": url, "error": str(e)},
|
||||
)
|
||||
return {"error": str(e)}
|
||||
@@ -13,6 +13,7 @@ class XAIUpstreamProvider(BaseUpstreamProvider):
|
||||
provider_type = "x-ai"
|
||||
default_base_url = "https://api.x.ai/v1"
|
||||
platform_url = "https://console.x.ai/"
|
||||
litellm_provider_prefix = "xai/"
|
||||
|
||||
def __init__(self, api_key: str, provider_fee: float = 1.01):
|
||||
super().__init__(
|
||||
|
||||
@@ -1,16 +1,32 @@
|
||||
import asyncio
|
||||
import math
|
||||
import time
|
||||
import typing
|
||||
from typing import TypedDict
|
||||
|
||||
from cashu.core.base import Proof, Token
|
||||
from cashu.core.mint_info import MintInfo as _CashuMintInfo
|
||||
from cashu.wallet.helpers import deserialize_token_from_string
|
||||
from cashu.wallet.wallet import Wallet
|
||||
from sqlmodel import col, update
|
||||
from pydantic_core import PydanticUndefined
|
||||
from sqlmodel import col, select, update
|
||||
|
||||
from .core import db, get_logger
|
||||
from .core.db import store_cashu_transaction
|
||||
from .core.settings import settings
|
||||
from .payment.lnurl import raw_send_to_lnurl
|
||||
|
||||
# cashu still declares Optional[X] without explicit defaults on MintInfo.
|
||||
# Under pydantic v2 those are required, but real mints omit many of them.
|
||||
# Default Optional fields to None at import time so balance fetches don't 422.
|
||||
for _name, _field in _CashuMintInfo.model_fields.items():
|
||||
_annot = _field.annotation
|
||||
_is_optional = typing.get_origin(_annot) is typing.Union and type(
|
||||
None
|
||||
) in typing.get_args(_annot)
|
||||
if _is_optional and _field.default is PydanticUndefined:
|
||||
_field.default = None
|
||||
_CashuMintInfo.model_rebuild(force=True)
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
@@ -32,24 +48,66 @@ async def recieve_token(
|
||||
if token_obj.mint not in settings.cashu_mints:
|
||||
return await swap_to_primary_mint(token_obj, wallet)
|
||||
|
||||
await wallet.load_mint(keyset_id=token_obj.keysets[0])
|
||||
|
||||
wallet.verify_proofs_dleq(token_obj.proofs)
|
||||
await wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True)
|
||||
|
||||
return token_obj.amount, token_obj.unit, token_obj.mint
|
||||
|
||||
|
||||
async def send(amount: int, unit: str, mint_url: str | None = None) -> tuple[int, str]:
|
||||
"""Internal send function - returns amount and serialized token"""
|
||||
wallet: Wallet = await get_wallet(mint_url or settings.primary_mint, unit)
|
||||
proofs = get_proofs_per_mint_and_unit(
|
||||
wallet, mint_url or settings.primary_mint, unit
|
||||
effective_mint_url = mint_url or settings.primary_mint
|
||||
wallet: Wallet = await get_wallet(effective_mint_url, unit)
|
||||
proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit)
|
||||
proofs_for_mint = sum(p.amount for p in proofs)
|
||||
|
||||
# Fallback: proofs from untrusted source mints are swapped to primary_mint
|
||||
# during receive, so the user's preferred refund_mint_url may have no proofs
|
||||
# even though the global wallet has the balance.
|
||||
if proofs_for_mint < amount and effective_mint_url != settings.primary_mint:
|
||||
logger.info(
|
||||
f"send: insufficient proofs at {effective_mint_url} "
|
||||
f"(have {proofs_for_mint}, need {amount}), falling back to primary_mint={settings.primary_mint}"
|
||||
)
|
||||
effective_mint_url = settings.primary_mint
|
||||
wallet = await get_wallet(effective_mint_url, unit)
|
||||
proofs = get_proofs_per_mint_and_unit(wallet, effective_mint_url, unit)
|
||||
proofs_for_mint = sum(p.amount for p in proofs)
|
||||
|
||||
all_mint_urls = list({k.mint_url for k in wallet.keysets.values()})
|
||||
proof_summary = {
|
||||
f"{k.mint_url}/{k.unit.name}": sum(p.amount for p in wallet.proofs if p.id == k.id)
|
||||
for k in wallet.keysets.values()
|
||||
}
|
||||
# Show ALL proofs in DB by keyset_id, regardless of whether the loaded wallet
|
||||
# knows about that keyset. This reveals proofs orphaned under stale keysets.
|
||||
raw_proofs_by_keyset: dict[str, int] = {}
|
||||
for p in wallet.proofs:
|
||||
raw_proofs_by_keyset[p.id] = raw_proofs_by_keyset.get(p.id, 0) + p.amount
|
||||
logger.info(
|
||||
f"send: proof inventory | mint={effective_mint_url} unit={unit} amount={amount} "
|
||||
f"primary_mint={settings.primary_mint} proofs_for_mint={proofs_for_mint} "
|
||||
f"all_mints={all_mint_urls} by_keyset={proof_summary} "
|
||||
f"raw_proofs_by_keyset_id={raw_proofs_by_keyset} "
|
||||
f"total_wallet_proofs={sum(p.amount for p in wallet.proofs)}"
|
||||
)
|
||||
|
||||
# Reserve proofs only after serialization succeeds — if serialize_proofs or
|
||||
# swap_to_send fails mid-way, proofs stay unreserved so dashboard balance
|
||||
# doesn't go negative.
|
||||
send_proofs, _ = await wallet.select_to_send(
|
||||
proofs, amount, set_reserved=True, include_fees=False
|
||||
)
|
||||
token = await wallet.serialize_proofs(
|
||||
send_proofs, include_dleq=False, legacy=False, memo=None
|
||||
proofs, amount, set_reserved=False, include_fees=False
|
||||
)
|
||||
try:
|
||||
token = await wallet.serialize_proofs(
|
||||
send_proofs, include_dleq=False, legacy=False, memo=None
|
||||
)
|
||||
except Exception:
|
||||
await wallet.set_reserved_for_send(send_proofs, reserved=False)
|
||||
raise
|
||||
await wallet.set_reserved_for_send(send_proofs, reserved=True)
|
||||
return amount, token
|
||||
|
||||
|
||||
@@ -58,15 +116,89 @@ async def send_token(amount: int, unit: str, mint_url: str | None = None) -> str
|
||||
return token
|
||||
|
||||
|
||||
async def _calculate_swap_amount(
|
||||
amount_msat: int,
|
||||
token_unit: str,
|
||||
token_mint_url: str,
|
||||
token_wallet: Wallet,
|
||||
primary_wallet: Wallet,
|
||||
proofs: list,
|
||||
) -> int:
|
||||
"""
|
||||
Calculate the amount to mint on the primary mint after accounting for
|
||||
melt fees and NUT-02 input fees on the foreign mint.
|
||||
"""
|
||||
if settings.primary_mint_unit == "sat":
|
||||
receive_amount = amount_msat // 1000
|
||||
else:
|
||||
receive_amount = amount_msat
|
||||
|
||||
if token_mint_url == settings.primary_mint:
|
||||
logger.info(
|
||||
"swap_to_primary_mint: skipping fee estimation (same mint)",
|
||||
extra={"minted_amount": receive_amount},
|
||||
)
|
||||
return int(receive_amount)
|
||||
|
||||
logger.info(
|
||||
"swap_to_primary_mint: estimating fees",
|
||||
extra={
|
||||
"dummy_amount": receive_amount,
|
||||
"unit": settings.primary_mint_unit,
|
||||
},
|
||||
)
|
||||
|
||||
try:
|
||||
dummy_mint_quote = await primary_wallet.request_mint(receive_amount)
|
||||
dummy_melt_quote = await token_wallet.melt_quote(dummy_mint_quote.request)
|
||||
|
||||
fee_reserve = dummy_melt_quote.fee_reserve
|
||||
input_fees = token_wallet.get_fees_for_proofs(proofs)
|
||||
if token_unit == "sat":
|
||||
fee_msat = (fee_reserve + input_fees) * 1000
|
||||
else:
|
||||
fee_msat = fee_reserve + input_fees
|
||||
|
||||
amount_msat_after_fee = amount_msat - fee_msat
|
||||
|
||||
if settings.primary_mint_unit == "sat":
|
||||
minted_amount = int(amount_msat_after_fee // 1000)
|
||||
else:
|
||||
minted_amount = int(amount_msat_after_fee)
|
||||
|
||||
if minted_amount <= 0:
|
||||
raise ValueError(f"Fees ({fee_reserve + input_fees} {token_unit}) exceed token amount")
|
||||
|
||||
logger.info(
|
||||
"swap_to_primary_mint: fee estimation result",
|
||||
extra={
|
||||
"token_amount_sat": amount_msat // 1000,
|
||||
"estimated_fee_sat": fee_msat // 1000,
|
||||
"input_fees": input_fees,
|
||||
"minted_amount": minted_amount,
|
||||
"minted_unit": settings.primary_mint_unit,
|
||||
},
|
||||
)
|
||||
return minted_amount
|
||||
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"swap_to_primary_mint: fee estimation failed",
|
||||
extra={"error": str(e)},
|
||||
)
|
||||
raise ValueError(f"Failed to estimate fees: {e}") from e
|
||||
|
||||
|
||||
async def swap_to_primary_mint(
|
||||
token_obj: Token, token_wallet: Wallet
|
||||
) -> tuple[int, str, str]:
|
||||
logger.info(
|
||||
"swap_to_primary_mint",
|
||||
"swap_to_primary_mint: starting",
|
||||
extra={
|
||||
"mint": token_obj.mint,
|
||||
"amount": token_obj.amount,
|
||||
"foreign_mint": token_obj.mint,
|
||||
"token_amount": token_obj.amount,
|
||||
"unit": token_obj.unit,
|
||||
"primary_mint": settings.primary_mint,
|
||||
},
|
||||
)
|
||||
# Ensure amount is an integer
|
||||
@@ -81,24 +213,169 @@ async def swap_to_primary_mint(
|
||||
amount_msat = token_amount
|
||||
else:
|
||||
raise ValueError("Invalid unit")
|
||||
estimated_fee_sat = math.ceil(max(amount_msat // 1000 * 0.01, 2)) + 1
|
||||
amount_msat_after_fee = amount_msat - estimated_fee_sat * 1000
|
||||
primary_wallet = await get_wallet(settings.primary_mint, settings.primary_mint_unit)
|
||||
|
||||
if settings.primary_mint_unit == "sat":
|
||||
minted_amount = int(amount_msat_after_fee // 1000)
|
||||
else:
|
||||
minted_amount = int(amount_msat_after_fee)
|
||||
# If the token is already from the primary mint, we don't need to swap
|
||||
# and we definitely don't want to calculate or pay fees.
|
||||
if token_obj.mint == settings.primary_mint:
|
||||
logger.info(
|
||||
"swap_to_primary_mint: token already on primary mint, skipping swap",
|
||||
extra={
|
||||
"mint": token_obj.mint,
|
||||
"amount": token_amount,
|
||||
"unit": token_obj.unit,
|
||||
},
|
||||
)
|
||||
await token_wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True)
|
||||
return token_amount, token_obj.unit, token_obj.mint
|
||||
|
||||
minted_amount = await _calculate_swap_amount(
|
||||
amount_msat,
|
||||
token_obj.unit,
|
||||
token_obj.mint,
|
||||
token_wallet,
|
||||
primary_wallet,
|
||||
token_obj.proofs,
|
||||
)
|
||||
|
||||
mint_quote = await primary_wallet.request_mint(minted_amount)
|
||||
logger.info(
|
||||
"swap_to_primary_mint: mint quote received",
|
||||
extra={"mint_quote_id": mint_quote.quote},
|
||||
)
|
||||
|
||||
melt_quote = await token_wallet.melt_quote(mint_quote.request)
|
||||
_ = await token_wallet.melt(
|
||||
proofs=token_obj.proofs,
|
||||
invoice=mint_quote.request,
|
||||
fee_reserve_sat=melt_quote.fee_reserve,
|
||||
quote_id=melt_quote.quote,
|
||||
input_fees = token_wallet.get_fees_for_proofs(token_obj.proofs)
|
||||
total_needed = melt_quote.amount + melt_quote.fee_reserve + input_fees
|
||||
logger.info(
|
||||
"swap_to_primary_mint: melt quote received",
|
||||
extra={
|
||||
"melt_quote_id": melt_quote.quote,
|
||||
"melt_amount": melt_quote.amount,
|
||||
"melt_fee_reserve": melt_quote.fee_reserve,
|
||||
"input_fees": input_fees,
|
||||
"total_needed": total_needed,
|
||||
"token_amount": token_amount,
|
||||
},
|
||||
)
|
||||
|
||||
if total_needed > token_amount:
|
||||
logger.warning(
|
||||
"swap_to_primary_mint: insufficient token amount for melt fees",
|
||||
extra={
|
||||
"token_amount": token_amount,
|
||||
"melt_amount": melt_quote.amount,
|
||||
"melt_fee_reserve": melt_quote.fee_reserve,
|
||||
"input_fees": input_fees,
|
||||
"total_needed": total_needed,
|
||||
"shortfall": total_needed - token_amount,
|
||||
},
|
||||
)
|
||||
raise ValueError(
|
||||
f"Token amount ({token_amount} {token_obj.unit}) is insufficient to cover "
|
||||
f"melt fees. Needed: {total_needed} {token_obj.unit} "
|
||||
f"(amount: {melt_quote.amount} + fee: {melt_quote.fee_reserve} + input_fees: {input_fees})"
|
||||
)
|
||||
|
||||
try:
|
||||
_ = await token_wallet.melt(
|
||||
proofs=token_obj.proofs,
|
||||
invoice=mint_quote.request,
|
||||
fee_reserve_sat=melt_quote.fee_reserve,
|
||||
quote_id=melt_quote.quote,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"swap_to_primary_mint: melt failed",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
"foreign_mint": token_obj.mint,
|
||||
"token_amount": token_amount,
|
||||
"melt_quote_id": melt_quote.quote,
|
||||
"total_needed": total_needed,
|
||||
},
|
||||
)
|
||||
raise ValueError(
|
||||
f"Failed to melt token from foreign mint {token_obj.mint}: {e}"
|
||||
) from e
|
||||
|
||||
logger.info(
|
||||
"swap_to_primary_mint: melt succeeded, minting on primary",
|
||||
extra={"minted_amount": minted_amount, "mint_quote_id": mint_quote.quote},
|
||||
)
|
||||
|
||||
await primary_wallet.load_proofs(reload=True)
|
||||
pre_mint_balance = primary_wallet.available_balance.amount
|
||||
try:
|
||||
_ = await primary_wallet.mint(minted_amount, quote_id=mint_quote.quote)
|
||||
except Exception as e:
|
||||
if "11003" in str(e) or "outputs already signed" in str(e).lower():
|
||||
# Previous mint call signed outputs at the mint but failed before
|
||||
# bump_secret_derivation ran locally. Recover orphaned proofs and
|
||||
# advance the counter so the next request derives fresh secrets.
|
||||
logger.warning(
|
||||
"swap_to_primary_mint: outputs already signed — recovering orphaned proofs",
|
||||
extra={"mint_quote_id": mint_quote.quote, "minted_amount": minted_amount},
|
||||
)
|
||||
try:
|
||||
for keyset_id in primary_wallet.keysets:
|
||||
await primary_wallet.restore_tokens_for_keyset(keyset_id, to=1, batch=25)
|
||||
await primary_wallet.load_proofs(reload=True)
|
||||
post_recovery_balance = primary_wallet.available_balance.amount
|
||||
balance_gained = post_recovery_balance - pre_mint_balance
|
||||
logger.info(
|
||||
"swap_to_primary_mint: recovery scan completed",
|
||||
extra={
|
||||
"pre_mint_balance": pre_mint_balance,
|
||||
"post_recovery_balance": post_recovery_balance,
|
||||
"balance_gained": balance_gained,
|
||||
"expected": minted_amount,
|
||||
},
|
||||
)
|
||||
if balance_gained < minted_amount:
|
||||
# Recovery scan ran but did NOT restore the orphaned proofs
|
||||
# (mint reports them as spent — they're stuck). Refuse to
|
||||
# credit the API key balance for proofs we don't actually hold.
|
||||
raise ValueError(
|
||||
f"Swap recovery failed: mint signed outputs but proofs are "
|
||||
f"unrecoverable (mint reports them spent). "
|
||||
f"Expected {minted_amount}, recovered {balance_gained}. "
|
||||
f"Local wallet DB ('.wallet/') state is corrupted — "
|
||||
f"the counter for keyset is stuck at a bad index range."
|
||||
)
|
||||
except ValueError:
|
||||
raise
|
||||
except Exception as recovery_err:
|
||||
logger.error(
|
||||
"swap_to_primary_mint: recovery failed",
|
||||
extra={"error": str(recovery_err)},
|
||||
)
|
||||
raise ValueError(
|
||||
f"Mint on primary failed and recovery unsuccessful: {e}"
|
||||
) from e
|
||||
else:
|
||||
logger.error(
|
||||
"swap_to_primary_mint: mint on primary failed after successful melt",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"error_type": type(e).__name__,
|
||||
"minted_amount": minted_amount,
|
||||
"mint_quote_id": mint_quote.quote,
|
||||
},
|
||||
)
|
||||
raise
|
||||
|
||||
logger.info(
|
||||
"swap_to_primary_mint: completed successfully",
|
||||
extra={
|
||||
"foreign_mint": token_obj.mint,
|
||||
"primary_mint": settings.primary_mint,
|
||||
"original_amount": token_amount,
|
||||
"minted_amount": minted_amount,
|
||||
"unit": settings.primary_mint_unit,
|
||||
},
|
||||
)
|
||||
_ = await primary_wallet.mint(minted_amount, quote_id=mint_quote.quote)
|
||||
|
||||
return int(minted_amount), settings.primary_mint_unit, settings.primary_mint
|
||||
|
||||
@@ -113,6 +390,8 @@ async def credit_balance(
|
||||
|
||||
try:
|
||||
amount, unit, mint_url = await recieve_token(cashu_token)
|
||||
original_amount = amount
|
||||
original_unit = unit
|
||||
logger.info(
|
||||
"credit_balance: Token redeemed successfully",
|
||||
extra={"amount": amount, "unit": unit, "mint_url": mint_url},
|
||||
@@ -144,7 +423,20 @@ async def credit_balance(
|
||||
extra={"new_balance": key.balance},
|
||||
)
|
||||
|
||||
logger.info(
|
||||
try:
|
||||
await store_cashu_transaction(
|
||||
token=cashu_token,
|
||||
amount=original_amount,
|
||||
unit=original_unit,
|
||||
mint_url=mint_url,
|
||||
typ="in",
|
||||
source="apikey",
|
||||
api_key_hashed_key=key.hashed_key,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
logger.debug(
|
||||
"Cashu token successfully redeemed and stored",
|
||||
extra={"amount": amount, "unit": unit, "mint_url": mint_url},
|
||||
)
|
||||
@@ -246,7 +538,7 @@ async def fetch_all_balances(
|
||||
"unit": unit,
|
||||
"wallet_balance": proofs_balance,
|
||||
"user_balance": user_balance,
|
||||
"owner_balance": proofs_balance - user_balance,
|
||||
"owner_balance": proofs_balance - user_balance if proofs_balance != 0 else 0,
|
||||
}
|
||||
return result
|
||||
except Exception as e:
|
||||
@@ -294,7 +586,9 @@ async def fetch_all_balances(
|
||||
total_wallet_balance_sats += proofs_balance_sats
|
||||
total_user_balance_sats += user_balance_sats
|
||||
|
||||
owner_balance = total_wallet_balance_sats - total_user_balance_sats
|
||||
owner_balance = 0
|
||||
if total_wallet_balance_sats != 0:
|
||||
owner_balance = total_wallet_balance_sats - total_user_balance_sats
|
||||
|
||||
return (
|
||||
balance_details,
|
||||
@@ -305,11 +599,11 @@ async def fetch_all_balances(
|
||||
|
||||
|
||||
async def periodic_payout() -> None:
|
||||
if not settings.receive_ln_address:
|
||||
logger.error("RECEIVE_LN_ADDRESS is not set, skipping payout")
|
||||
return
|
||||
while True:
|
||||
await asyncio.sleep(60 * 15)
|
||||
await asyncio.sleep(settings.payout_interval_seconds)
|
||||
print(settings.payout_interval_seconds)
|
||||
if not settings.receive_ln_address:
|
||||
continue
|
||||
try:
|
||||
async with db.create_session() as session:
|
||||
for mint_url in settings.cashu_mints:
|
||||
@@ -319,6 +613,7 @@ async def periodic_payout() -> None:
|
||||
wallet, mint_url, unit, not_reserved=True
|
||||
)
|
||||
proofs = await slow_filter_spend_proofs(proofs, wallet)
|
||||
await asyncio.sleep(5)
|
||||
user_balance = await db.balances_for_mint_and_unit(
|
||||
session, mint_url, unit
|
||||
)
|
||||
@@ -326,7 +621,12 @@ async def periodic_payout() -> None:
|
||||
user_balance = user_balance // 1000
|
||||
proofs_balance = sum(proof.amount for proof in proofs)
|
||||
available_balance = proofs_balance - user_balance
|
||||
min_amount = 210 if unit == "sat" else 210000
|
||||
# Threshold is configured in sats; convert for msat wallets.
|
||||
min_amount = (
|
||||
settings.min_payout_sat
|
||||
if unit == "sat"
|
||||
else settings.min_payout_sat * 1000
|
||||
)
|
||||
if available_balance > min_amount:
|
||||
amount_received = await raw_send_to_lnurl(
|
||||
wallet,
|
||||
@@ -344,8 +644,6 @@ async def periodic_payout() -> None:
|
||||
"amount_received": amount_received,
|
||||
},
|
||||
)
|
||||
|
||||
await asyncio.sleep(5)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Error sending payout: {type(e).__name__}",
|
||||
@@ -353,6 +651,101 @@ async def periodic_payout() -> None:
|
||||
)
|
||||
|
||||
|
||||
async def periodic_refund_sweep() -> None:
|
||||
while True:
|
||||
await asyncio.sleep(60 * 60) # every hour
|
||||
try:
|
||||
cutoff = int(time.time()) - settings.refund_sweep_ttl_seconds
|
||||
async with db.create_session() as session:
|
||||
stmt = select(db.CashuTransaction).where(
|
||||
db.CashuTransaction.type == "out",
|
||||
db.CashuTransaction.collected == False, # noqa: E712
|
||||
db.CashuTransaction.swept == False, # noqa: E712
|
||||
db.CashuTransaction.created_at < cutoff,
|
||||
)
|
||||
results = await session.exec(stmt)
|
||||
refunds = results.all()
|
||||
|
||||
for refund in refunds:
|
||||
try:
|
||||
await recieve_token(refund.token)
|
||||
refund.swept = True
|
||||
session.add(refund)
|
||||
logger.info(
|
||||
"Swept uncollected refund",
|
||||
extra={
|
||||
"id": refund.id,
|
||||
"amount": refund.amount,
|
||||
"unit": refund.unit,
|
||||
},
|
||||
)
|
||||
except Exception as e:
|
||||
error_msg = str(e).lower()
|
||||
if "already spent" in error_msg:
|
||||
refund.collected = True
|
||||
session.add(refund)
|
||||
logger.info(
|
||||
"Refund already spent (client collected), marking swept",
|
||||
extra={
|
||||
"id": refund.id,
|
||||
},
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"Failed to sweep refund",
|
||||
extra={
|
||||
"id": refund.id,
|
||||
"error": str(e),
|
||||
},
|
||||
)
|
||||
await session.commit()
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Error in periodic refund sweep",
|
||||
extra={"error": str(e), "error_type": type(e).__name__},
|
||||
)
|
||||
|
||||
|
||||
async def periodic_routstr_fee_payout() -> None:
|
||||
from .auth import (
|
||||
ROUTSTR_FEE_DEFAULT_PAYOUT,
|
||||
ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS,
|
||||
ROUTSTR_LN_ADDRESS,
|
||||
)
|
||||
|
||||
if not ROUTSTR_LN_ADDRESS:
|
||||
logger.info("ROUTSTR_LN_ADDRESS not set, skipping fee payout")
|
||||
return
|
||||
while True:
|
||||
await asyncio.sleep(ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS)
|
||||
try:
|
||||
async with db.create_session() as session:
|
||||
fee = await db.get_routstr_fee(session)
|
||||
accumulated_sats = fee.accumulated_msats // 1000
|
||||
if accumulated_sats >= ROUTSTR_FEE_DEFAULT_PAYOUT:
|
||||
wallet = await get_wallet(settings.primary_mint, "sat")
|
||||
proofs = get_proofs_per_mint_and_unit(
|
||||
wallet, settings.primary_mint, "sat", not_reserved=True
|
||||
)
|
||||
amount_received = await raw_send_to_lnurl(
|
||||
wallet, proofs, ROUTSTR_LN_ADDRESS, "sat", amount=accumulated_sats
|
||||
)
|
||||
paid_msats = accumulated_sats * 1000
|
||||
await db.reset_routstr_fee(session, paid_msats)
|
||||
logger.info(
|
||||
"Routstr fee payout sent",
|
||||
extra={
|
||||
"accumulated_sats": accumulated_sats,
|
||||
"amount_received": amount_received,
|
||||
},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Error in Routstr fee payout: {type(e).__name__}",
|
||||
extra={"error": str(e)},
|
||||
)
|
||||
|
||||
|
||||
async def send_to_lnurl(amount: int, unit: str, mint: str, address: str) -> int:
|
||||
wallet = await get_wallet(mint, unit)
|
||||
proofs = wallet._get_proofs_per_keyset(wallet.proofs)[wallet.keyset_id]
|
||||
|
||||
@@ -380,6 +380,13 @@ async def integration_session(
|
||||
yield session
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def patched_db_engine(integration_engine: Any) -> AsyncGenerator[None, None]:
|
||||
"""Patch the global db engine so create_session() uses the test engine."""
|
||||
with patch("routstr.core.db.engine", integration_engine):
|
||||
yield
|
||||
|
||||
|
||||
class DatabaseSnapshot:
|
||||
"""Utility to capture and compare database states"""
|
||||
|
||||
|
||||
80
tests/integration/test_admin_provider_balance.py
Normal file
80
tests/integration/test_admin_provider_balance.py
Normal file
@@ -0,0 +1,80 @@
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.admin import admin_sessions
|
||||
from routstr.core.db import UpstreamProviderRow
|
||||
|
||||
|
||||
async def _create_routstr_provider() -> UpstreamProviderRow:
|
||||
return UpstreamProviderRow(
|
||||
provider_type="routstr",
|
||||
base_url="https://upstream.example",
|
||||
api_key="",
|
||||
enabled=True,
|
||||
)
|
||||
|
||||
|
||||
def _admin_headers() -> dict[str, str]:
|
||||
token = "test-admin-token"
|
||||
admin_sessions[token] = int(
|
||||
(datetime.now(timezone.utc) + timedelta(minutes=5)).timestamp()
|
||||
)
|
||||
return {"Authorization": f"Bearer {token}"}
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_admin_routstr_provider_balance_timeout_returns_504(
|
||||
integration_client: httpx.AsyncClient,
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
provider = await _create_routstr_provider()
|
||||
integration_session.add(provider)
|
||||
await integration_session.commit()
|
||||
await integration_session.refresh(provider)
|
||||
|
||||
request = httpx.Request("GET", f"{provider.base_url}/v1/balance/info")
|
||||
timeout_error = httpx.ConnectTimeout("Connect timeout", request=request)
|
||||
|
||||
with patch(
|
||||
"httpx.AsyncHTTPTransport.handle_async_request",
|
||||
new=AsyncMock(side_effect=timeout_error),
|
||||
):
|
||||
response = await integration_client.get(
|
||||
f"/admin/api/upstream-providers/{provider.id}/balance",
|
||||
headers=_admin_headers(),
|
||||
)
|
||||
|
||||
assert response.status_code == 504
|
||||
assert response.json()["detail"] == "Timed out contacting upstream Routstr provider"
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_admin_routstr_provider_balance_request_error_returns_502(
|
||||
integration_client: httpx.AsyncClient,
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
provider = await _create_routstr_provider()
|
||||
integration_session.add(provider)
|
||||
await integration_session.commit()
|
||||
await integration_session.refresh(provider)
|
||||
|
||||
request = httpx.Request("GET", f"{provider.base_url}/v1/balance/info")
|
||||
request_error = httpx.ConnectError("Connection failed", request=request)
|
||||
|
||||
with patch(
|
||||
"httpx.AsyncHTTPTransport.handle_async_request",
|
||||
new=AsyncMock(side_effect=request_error),
|
||||
):
|
||||
response = await integration_client.get(
|
||||
f"/admin/api/upstream-providers/{provider.id}/balance",
|
||||
headers=_admin_headers(),
|
||||
)
|
||||
|
||||
assert response.status_code == 502
|
||||
assert response.json()["detail"] == "Failed to contact upstream Routstr provider"
|
||||
369
tests/integration/test_balance_negative_on_cost_overrun.py
Normal file
369
tests/integration/test_balance_negative_on_cost_overrun.py
Normal file
@@ -0,0 +1,369 @@
|
||||
"""
|
||||
Integration tests for the balance-goes-negative bug in adjust_payment_for_tokens.
|
||||
|
||||
Root cause: when actual token cost exceeds the discounted reservation
|
||||
(cost_difference > 0, caused by tolerance_percentage discounting the reservation),
|
||||
the finalization UPDATE had no WHERE guard on balance, allowing balance to go negative.
|
||||
|
||||
Fix: added `.where(col(ApiKey.balance) >= total_cost_msats)` so the UPDATE is a no-op
|
||||
when balance is insufficient, then falls back to charging only deducted_max_cost.
|
||||
"""
|
||||
|
||||
import uuid
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.db import ApiKey
|
||||
from routstr.payment.cost_calculation import CostData
|
||||
|
||||
|
||||
def _make_key(balance: int, reserved: int) -> ApiKey:
|
||||
return ApiKey(
|
||||
hashed_key=f"test_{uuid.uuid4().hex}",
|
||||
balance=balance,
|
||||
reserved_balance=reserved,
|
||||
total_spent=0,
|
||||
total_requests=1,
|
||||
)
|
||||
|
||||
|
||||
async def _refresh(session: AsyncSession, key: ApiKey) -> ApiKey:
|
||||
await session.refresh(key)
|
||||
return key
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helper: build a CostData where token cost > deducted_max_cost
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _cost_data(total_msats: int) -> CostData:
|
||||
return CostData(
|
||||
base_msats=0,
|
||||
input_msats=total_msats // 2,
|
||||
output_msats=total_msats - total_msats // 2,
|
||||
total_msats=total_msats,
|
||||
total_usd=0.0,
|
||||
input_tokens=100,
|
||||
output_tokens=100,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 1 — exact reproduction of the bug
|
||||
#
|
||||
# Setup: balance == deducted_max_cost (user has just enough for the reservation,
|
||||
# nothing extra). Actual token cost is 1% higher (tolerance_percentage).
|
||||
#
|
||||
# Before fix: balance -= total_cost_msats → goes negative.
|
||||
# After fix: WHERE balance >= total_cost_msats fails → fallback charges
|
||||
# deducted_max_cost → balance reaches 0, never negative.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_balance_never_negative_when_cost_exceeds_reservation(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Balance must not go negative when actual token cost > discounted reservation."""
|
||||
from routstr.auth import adjust_payment_for_tokens
|
||||
|
||||
deducted_max_cost = 990 # reserved (1% below true max of 1000)
|
||||
actual_token_cost = 1000 # actual cost at true max
|
||||
|
||||
# User has balance exactly equal to the reservation — tight budget
|
||||
key = _make_key(balance=deducted_max_cost, reserved=deducted_max_cost)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
response_data = {"model": "test-model", "usage": {"prompt_tokens": 100, "completion_tokens": 100}}
|
||||
|
||||
with patch(
|
||||
"routstr.auth.calculate_cost",
|
||||
return_value=_cost_data(actual_token_cost),
|
||||
):
|
||||
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
|
||||
|
||||
await _refresh(integration_session, key)
|
||||
|
||||
assert key.balance >= 0, f"Balance went negative: {key.balance}"
|
||||
assert key.reserved_balance >= 0, f"Reserved balance went negative: {key.reserved_balance}"
|
||||
assert key.reserved_balance == 0, "Reservation must be fully released after finalization"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 2 — balance is ZERO after the reservation is accounted for
|
||||
# (absolute floor case)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_balance_floor_at_zero_on_overrun(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""When balance exactly covers deducted_max_cost and cost overruns, balance reaches 0 not negative."""
|
||||
from routstr.auth import adjust_payment_for_tokens
|
||||
|
||||
deducted_max_cost = 500
|
||||
actual_token_cost = 550 # 10% overrun
|
||||
|
||||
key = _make_key(balance=500, reserved=500)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 50}}
|
||||
|
||||
with patch(
|
||||
"routstr.auth.calculate_cost",
|
||||
return_value=_cost_data(actual_token_cost),
|
||||
):
|
||||
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
|
||||
|
||||
await _refresh(integration_session, key)
|
||||
|
||||
assert key.balance == 0, (
|
||||
f"Expected balance=0 (charged deducted_max_cost fallback), got {key.balance}"
|
||||
)
|
||||
assert key.reserved_balance == 0, f"Reserved balance should be 0, got {key.reserved_balance}"
|
||||
# Fallback charges deducted_max_cost
|
||||
assert key.total_spent == deducted_max_cost, (
|
||||
f"Expected total_spent={deducted_max_cost}, got {key.total_spent}"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 3 — balance has enough room: full token cost should be charged
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_full_cost_charged_when_balance_sufficient_for_overrun(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""When balance covers total_cost_msats, the full amount is charged (not just deducted_max_cost)."""
|
||||
from routstr.auth import adjust_payment_for_tokens
|
||||
|
||||
deducted_max_cost = 990
|
||||
actual_token_cost = 1000
|
||||
|
||||
# User has extra balance beyond the reservation
|
||||
key = _make_key(balance=2000, reserved=990)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
response_data = {"model": "test-model", "usage": {"prompt_tokens": 100, "completion_tokens": 100}}
|
||||
|
||||
with patch(
|
||||
"routstr.auth.calculate_cost",
|
||||
return_value=_cost_data(actual_token_cost),
|
||||
):
|
||||
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
|
||||
|
||||
await _refresh(integration_session, key)
|
||||
|
||||
assert key.balance >= 0, f"Balance went negative: {key.balance}"
|
||||
assert key.reserved_balance == 0, f"Reservation not released: {key.reserved_balance}"
|
||||
assert key.total_spent == actual_token_cost, (
|
||||
f"Expected full charge of {actual_token_cost}, got {key.total_spent}"
|
||||
)
|
||||
assert key.balance == 2000 - actual_token_cost, (
|
||||
f"Expected balance={2000 - actual_token_cost}, got {key.balance}"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 4 — concurrent finalizations with cost overrun
|
||||
#
|
||||
# Multiple requests finish concurrently. Each has a small overrun.
|
||||
# None should drive balance negative.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_cost_overruns_never_negative(
|
||||
integration_session: AsyncSession,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
"""Concurrent finalization with cost overruns must never produce negative balance."""
|
||||
import asyncio
|
||||
|
||||
from routstr.auth import adjust_payment_for_tokens, pay_for_request
|
||||
from routstr.core.db import create_session
|
||||
|
||||
deducted_max_cost = 990
|
||||
actual_token_cost = 1000
|
||||
n_requests = 5
|
||||
|
||||
# Fund the key with exactly enough for n_requests reservations + a tiny buffer
|
||||
starting_balance = deducted_max_cost * n_requests
|
||||
key_hash = f"test_concurrent_{uuid.uuid4().hex}"
|
||||
|
||||
async with create_session() as session:
|
||||
key = ApiKey(
|
||||
hashed_key=key_hash,
|
||||
balance=starting_balance,
|
||||
reserved_balance=0,
|
||||
total_spent=0,
|
||||
total_requests=0,
|
||||
)
|
||||
session.add(key)
|
||||
await session.commit()
|
||||
|
||||
# Reserve n_requests slots (sequentially, as pay_for_request is atomic)
|
||||
async with create_session() as session:
|
||||
key_to_reserve = await session.get(ApiKey, key_hash)
|
||||
assert key_to_reserve is not None
|
||||
for _ in range(n_requests):
|
||||
await pay_for_request(key_to_reserve, deducted_max_cost, session)
|
||||
await session.refresh(key_to_reserve)
|
||||
|
||||
# Now finalize all concurrently with cost overrun
|
||||
async def finalize() -> None:
|
||||
response_data = {
|
||||
"model": "test-model",
|
||||
"usage": {"prompt_tokens": 100, "completion_tokens": 100},
|
||||
}
|
||||
async with create_session() as session:
|
||||
fresh_key = await session.get(ApiKey, key_hash)
|
||||
assert fresh_key is not None
|
||||
with patch(
|
||||
"routstr.auth.calculate_cost",
|
||||
return_value=_cost_data(actual_token_cost),
|
||||
):
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key, response_data, session, deducted_max_cost
|
||||
)
|
||||
|
||||
await asyncio.gather(*[finalize() for _ in range(n_requests)])
|
||||
|
||||
async with create_session() as session:
|
||||
final_key = await session.get(ApiKey, key_hash)
|
||||
assert final_key is not None
|
||||
|
||||
assert final_key.balance >= 0, (
|
||||
f"Balance went negative after concurrent overruns: {final_key.balance}"
|
||||
)
|
||||
assert final_key.reserved_balance == 0, (
|
||||
f"Reserved balance not fully released: {final_key.reserved_balance}"
|
||||
)
|
||||
assert final_key.total_spent <= starting_balance, (
|
||||
f"Total spent ({final_key.total_spent}) exceeds starting balance ({starting_balance})"
|
||||
)
|
||||
# Every request must have been charged at least deducted_max_cost — no free inference.
|
||||
assert final_key.total_spent == starting_balance, (
|
||||
f"Expected total_spent={starting_balance} (all {n_requests} reservations charged), "
|
||||
f"got {final_key.total_spent} — at least one request got free inference"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 5 — overrun with no balance at all (reserved_balance == balance)
|
||||
# simulates a user who topped up to exactly the reservation floor
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_zero_free_balance_overrun_is_safe(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""User with zero free balance (all reserved) should never go negative on overrun."""
|
||||
from routstr.auth import adjust_payment_for_tokens
|
||||
|
||||
deducted_max_cost = 1000
|
||||
actual_token_cost = 1050
|
||||
|
||||
# balance == reserved_balance: zero free balance
|
||||
key = _make_key(balance=1000, reserved=1000)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 100}}
|
||||
|
||||
with patch(
|
||||
"routstr.auth.calculate_cost",
|
||||
return_value=_cost_data(actual_token_cost),
|
||||
):
|
||||
await adjust_payment_for_tokens(key, response_data, integration_session, deducted_max_cost)
|
||||
|
||||
await _refresh(integration_session, key)
|
||||
|
||||
assert key.balance >= 0, f"Balance went negative: {key.balance}"
|
||||
assert key.reserved_balance >= 0, f"Reserved balance went negative: {key.reserved_balance}"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 6 — parallel requests: second finalization must not get free inference
|
||||
#
|
||||
# Root cause of the bug fixed in auth.py:
|
||||
# `.where(col(ApiKey.balance) >= total_cost_msats)` ignores other requests'
|
||||
# reservations, so after Request A charges total_cost_msats, balance can drop
|
||||
# below deducted_max_cost, causing Request B's fallback to release for free.
|
||||
#
|
||||
# Fix: use `balance - reserved_balance + deducted_max_cost >= total_cost_msats`
|
||||
# so the check accounts for concurrent reservations.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_parallel_requests_no_free_inference(
|
||||
integration_session: AsyncSession,
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
"""Second parallel finalization must be charged even when first depleted free balance."""
|
||||
import asyncio
|
||||
|
||||
from routstr.auth import adjust_payment_for_tokens
|
||||
from routstr.core.db import create_session
|
||||
|
||||
deducted_max_cost = 100
|
||||
actual_token_cost = 150 # overrun: 50 more than reserved
|
||||
|
||||
# Fund the key with exactly 2 * deducted_max_cost.
|
||||
# Both requests pre-reserved 100 each → balance=200, reserved=200, free=0.
|
||||
# Old check (balance >= total_cost_msats):
|
||||
# Request A: 200 >= 150 ✓ → charges 150 → balance=50, reserved=100
|
||||
# Request B: 50 >= 150 ✗ → fallback: 50 >= 100 ✗ → releases FREE
|
||||
# New check (balance - reserved + deducted >= total_cost_msats):
|
||||
# Both fall to fallback (0 free balance).
|
||||
# Both charge deducted_max_cost=100 → total_spent=200, balance=0.
|
||||
starting_balance = deducted_max_cost * 2
|
||||
key_hash = f"test_parallel_no_free_{uuid.uuid4().hex}"
|
||||
|
||||
async with create_session() as session:
|
||||
key = ApiKey(
|
||||
hashed_key=key_hash,
|
||||
balance=starting_balance,
|
||||
reserved_balance=deducted_max_cost * 2, # both slots pre-reserved
|
||||
total_spent=0,
|
||||
total_requests=2,
|
||||
)
|
||||
session.add(key)
|
||||
await session.commit()
|
||||
|
||||
async def finalize() -> None:
|
||||
response_data = {
|
||||
"model": "test-model",
|
||||
"usage": {"prompt_tokens": 50, "completion_tokens": 100},
|
||||
}
|
||||
async with create_session() as session:
|
||||
fresh_key = await session.get(ApiKey, key_hash)
|
||||
assert fresh_key is not None
|
||||
with patch(
|
||||
"routstr.auth.calculate_cost",
|
||||
return_value=_cost_data(actual_token_cost),
|
||||
):
|
||||
await adjust_payment_for_tokens(
|
||||
fresh_key, response_data, session, deducted_max_cost
|
||||
)
|
||||
|
||||
await asyncio.gather(finalize(), finalize())
|
||||
|
||||
async with create_session() as session:
|
||||
final_key = await session.get(ApiKey, key_hash)
|
||||
assert final_key is not None
|
||||
|
||||
assert final_key.balance >= 0, f"Balance went negative: {final_key.balance}"
|
||||
assert final_key.reserved_balance == 0, (
|
||||
f"Reserved balance not released: {final_key.reserved_balance}"
|
||||
)
|
||||
# Both requests must have been charged — no free inference.
|
||||
assert final_key.total_spent == starting_balance, (
|
||||
f"Expected total_spent={starting_balance} (both reservations charged), "
|
||||
f"got {final_key.total_spent} — one request got free inference"
|
||||
)
|
||||
141
tests/integration/test_child_keys.py
Normal file
141
tests/integration/test_child_keys.py
Normal file
@@ -0,0 +1,141 @@
|
||||
import secrets
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.auth import adjust_payment_for_tokens, pay_for_request
|
||||
from routstr.balance import ChildKeyRequest, create_child_key
|
||||
from routstr.core.db import ApiKey
|
||||
from routstr.core.settings import settings
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_child_key_flow(integration_session: AsyncSession) -> None:
|
||||
# 1. Create a parent key with balance
|
||||
parent_raw = "parent_test_key_" + secrets.token_hex(4)
|
||||
parent_key = ApiKey(
|
||||
hashed_key=parent_raw,
|
||||
balance=10000, # 10 sats
|
||||
)
|
||||
integration_session.add(parent_key)
|
||||
await integration_session.commit()
|
||||
await integration_session.refresh(parent_key)
|
||||
|
||||
# Mock settings
|
||||
settings.child_key_cost = 1000 # 1 sat
|
||||
|
||||
# 2. Call create_child_key
|
||||
result = await create_child_key(
|
||||
ChildKeyRequest(count=1), parent_key, integration_session
|
||||
)
|
||||
|
||||
assert "api_keys" in result
|
||||
assert result["cost_msats"] == 1000
|
||||
assert result["parent_balance"] == 9000
|
||||
|
||||
child_key_raw = result["api_keys"][0][3:] # remove sk-
|
||||
|
||||
# 3. Verify child key exists in DB
|
||||
child_key_db = await integration_session.get(ApiKey, child_key_raw)
|
||||
assert child_key_db is not None
|
||||
assert child_key_db.parent_key_hash == parent_key.hashed_key
|
||||
assert child_key_db.balance == 0
|
||||
|
||||
# 4. Test payment with child key
|
||||
cost = 500
|
||||
await pay_for_request(child_key_db, cost, integration_session)
|
||||
|
||||
# Refresh keys
|
||||
await integration_session.refresh(parent_key)
|
||||
await integration_session.refresh(child_key_db)
|
||||
|
||||
# Parent should be charged
|
||||
assert parent_key.reserved_balance == 500
|
||||
assert parent_key.total_requests == 1
|
||||
|
||||
# Child should have total_requests incremented
|
||||
assert child_key_db.total_requests == 1
|
||||
|
||||
# 5. Test adjustment
|
||||
response_data = {"model": "test-model", "usage": {"total_tokens": 10}}
|
||||
|
||||
# Mock calculate_cost
|
||||
import routstr.auth
|
||||
from routstr.payment.cost_calculation import CostData
|
||||
|
||||
async def mock_calculate_cost(*args: Any, **kwargs: Any) -> CostData:
|
||||
return CostData(
|
||||
base_msats=0, input_msats=200, output_msats=200, total_msats=400
|
||||
)
|
||||
|
||||
# Patch calculate_cost
|
||||
original_calculate_cost = routstr.auth.calculate_cost
|
||||
routstr.auth.calculate_cost = mock_calculate_cost
|
||||
|
||||
try:
|
||||
adjustment = await adjust_payment_for_tokens(
|
||||
child_key_db, response_data, integration_session, 500
|
||||
)
|
||||
assert adjustment["total_msats"] == 400
|
||||
|
||||
# Refresh keys
|
||||
await integration_session.refresh(parent_key)
|
||||
await integration_session.refresh(child_key_db)
|
||||
|
||||
# Parent should have updated balance and total_spent
|
||||
assert parent_key.reserved_balance == 0
|
||||
assert parent_key.balance == 9000 - 400
|
||||
assert (
|
||||
parent_key.total_spent == 1400
|
||||
) # 1000 for child key creation + 400 for request
|
||||
|
||||
# Child should also have total_spent updated
|
||||
assert child_key_db.total_spent == 400
|
||||
|
||||
finally:
|
||||
routstr.auth.calculate_cost = original_calculate_cost
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_child_key_insufficient_balance(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
parent_key = ApiKey(
|
||||
hashed_key="poor_parent_" + secrets.token_hex(4),
|
||||
balance=500,
|
||||
)
|
||||
integration_session.add(parent_key)
|
||||
await integration_session.commit()
|
||||
await integration_session.refresh(parent_key)
|
||||
|
||||
settings.child_key_cost = 1000
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await create_child_key(
|
||||
ChildKeyRequest(count=1), parent_key, integration_session
|
||||
)
|
||||
assert exc.value.status_code == 402
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_child_key_cannot_create_child(integration_session: AsyncSession) -> None:
|
||||
parent_key = ApiKey(
|
||||
hashed_key="parent_" + secrets.token_hex(4),
|
||||
balance=10000,
|
||||
)
|
||||
child_key = ApiKey(
|
||||
hashed_key="child_" + secrets.token_hex(4),
|
||||
balance=0,
|
||||
parent_key_hash=parent_key.hashed_key,
|
||||
)
|
||||
integration_session.add(parent_key)
|
||||
integration_session.add(child_key)
|
||||
await integration_session.commit()
|
||||
await integration_session.refresh(child_key)
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await create_child_key(ChildKeyRequest(count=1), child_key, integration_session)
|
||||
assert exc.value.status_code == 400
|
||||
assert "Cannot create a child key for another child key" in str(exc.value.detail)
|
||||
97
tests/integration/test_child_keys_api.py
Normal file
97
tests/integration/test_child_keys_api.py
Normal file
@@ -0,0 +1,97 @@
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from httpx import AsyncClient
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_wallet_info_returns_child_keys(
|
||||
integration_client: AsyncClient,
|
||||
authenticated_client: AsyncClient,
|
||||
integration_session: Any,
|
||||
) -> None:
|
||||
"""Test that GET /v1/wallet/info returns child keys for a parent key"""
|
||||
|
||||
# 1. Get parent info to find its hashed_key
|
||||
response = await authenticated_client.get("/v1/wallet/info")
|
||||
assert response.status_code == 200
|
||||
parent_data = response.json()
|
||||
parent_data["api_key"]
|
||||
|
||||
# 2. Create child keys for this parent
|
||||
# We need to use the parent's authentication for this
|
||||
child_payload = {"count": 2, "balance_limit": 1000, "balance_limit_reset": "daily"}
|
||||
create_response = await authenticated_client.post(
|
||||
"/v1/wallet/child-key", json=child_payload
|
||||
)
|
||||
assert create_response.status_code == 200
|
||||
create_data = create_response.json()
|
||||
child_keys = create_data["api_keys"]
|
||||
assert len(child_keys) == 2
|
||||
|
||||
# 3. Call /info again and check for child_keys
|
||||
info_response = await authenticated_client.get("/v1/wallet/info")
|
||||
assert info_response.status_code == 200
|
||||
info_data = info_response.json()
|
||||
|
||||
assert "child_keys" in info_data
|
||||
assert len(info_data["child_keys"]) == 2
|
||||
|
||||
# Verify child key details
|
||||
for ck in info_data["child_keys"]:
|
||||
assert ck["api_key"] in child_keys
|
||||
assert ck["balance_limit"] == 1000
|
||||
assert ck["balance_limit_reset"] == "daily"
|
||||
assert "total_spent" in ck
|
||||
assert "total_requests" in ck
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_wallet_info_child_key_no_child_keys(
|
||||
integration_client: AsyncClient,
|
||||
authenticated_client: AsyncClient,
|
||||
integration_session: Any,
|
||||
) -> None:
|
||||
"""Test that GET /v1/wallet/info for a child key does NOT return child_keys"""
|
||||
|
||||
# 1. Create a child key
|
||||
child_payload = {"count": 1}
|
||||
create_response = await authenticated_client.post(
|
||||
"/v1/wallet/child-key", json=child_payload
|
||||
)
|
||||
assert create_response.status_code == 200
|
||||
child_key = create_response.json()["api_keys"][0]
|
||||
|
||||
# 2. Use the child key to get its info
|
||||
integration_client.headers["Authorization"] = f"Bearer {child_key}"
|
||||
info_response = await integration_client.get("/v1/wallet/info")
|
||||
assert info_response.status_code == 200
|
||||
info_data = info_response.json()
|
||||
|
||||
assert info_data["is_child"] is True
|
||||
assert "child_keys" not in info_data
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_account_info_root_returns_child_keys(
|
||||
authenticated_client: AsyncClient,
|
||||
) -> None:
|
||||
"""Test that GET / returns child keys for a parent key (root endpoint)"""
|
||||
|
||||
# 1. Create a child key
|
||||
child_payload = {"count": 1}
|
||||
await authenticated_client.post("/v1/wallet/child-key", json=child_payload)
|
||||
|
||||
# 2. Call root endpoint /v1/balance/
|
||||
# Note: routstr/balance.py defines router = APIRouter()
|
||||
# and it is included in balance_router with prefix /v1/balance
|
||||
# The endpoint is @router.get("/")
|
||||
response = await authenticated_client.get("/v1/balance/")
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
|
||||
assert "child_keys" in data
|
||||
assert len(data["child_keys"]) >= 1
|
||||
302
tests/integration/test_cli_tokens.py
Normal file
302
tests/integration/test_cli_tokens.py
Normal file
@@ -0,0 +1,302 @@
|
||||
"""Integration tests for CLI token management (/admin/api/cli-tokens).
|
||||
|
||||
Covers:
|
||||
- GET /admin/api/cli-tokens — list (preview only, no full token)
|
||||
- POST /admin/api/cli-tokens — create (returns full token once)
|
||||
- DELETE /admin/api/cli-tokens/{id} — revoke
|
||||
- Using a CLI token as Bearer auth against admin endpoints
|
||||
- Expiry enforcement (expired tokens are rejected by require_admin_api)
|
||||
- last_used_at bump on successful use
|
||||
- Auth failures: missing token, wrong token, revoked token
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import secrets
|
||||
import time
|
||||
from typing import AsyncGenerator
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from httpx import AsyncClient
|
||||
from sqlmodel import select
|
||||
|
||||
from routstr.core.admin import admin_sessions
|
||||
from routstr.core.db import AsyncSession, CliToken
|
||||
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
# Fixtures
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def admin_session_token() -> AsyncGenerator[str, None]:
|
||||
"""Inject a short-lived admin session token into admin_sessions."""
|
||||
token = secrets.token_urlsafe(24)
|
||||
admin_sessions[token] = int(time.time()) + 3600
|
||||
yield token
|
||||
admin_sessions.pop(token, None)
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def admin_client(
|
||||
integration_client: AsyncClient, admin_session_token: str
|
||||
) -> AsyncClient:
|
||||
"""An integration_client pre-authenticated with an admin session token."""
|
||||
integration_client.headers["Authorization"] = f"Bearer {admin_session_token}"
|
||||
return integration_client
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
# Creation
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_cli_token_returns_full_token_once(
|
||||
admin_client: AsyncClient,
|
||||
) -> None:
|
||||
"""POST /admin/api/cli-tokens returns the raw token only on creation."""
|
||||
resp = await admin_client.post(
|
||||
"/admin/api/cli-tokens",
|
||||
json={"name": "my-laptop"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
|
||||
assert body["name"] == "my-laptop"
|
||||
assert isinstance(body["id"], str) and body["id"]
|
||||
assert isinstance(body["token"], str) and len(body["token"]) >= 32
|
||||
assert body["expires_at"] is None
|
||||
assert isinstance(body["created_at"], int)
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_cli_token_with_expiry(admin_client: AsyncClient) -> None:
|
||||
"""expires_in_days sets expires_at ~= now + days * 86400."""
|
||||
before = int(time.time())
|
||||
resp = await admin_client.post(
|
||||
"/admin/api/cli-tokens",
|
||||
json={"name": "ci-runner", "expires_in_days": 7},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
|
||||
assert body["expires_at"] is not None
|
||||
delta = body["expires_at"] - before
|
||||
# Allow 10s jitter around 7 * 86400
|
||||
assert 7 * 86400 - 10 <= delta <= 7 * 86400 + 10
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_cli_token_rejects_empty_name(
|
||||
admin_client: AsyncClient,
|
||||
) -> None:
|
||||
resp = await admin_client.post(
|
||||
"/admin/api/cli-tokens", json={"name": " "}
|
||||
)
|
||||
assert resp.status_code == 400
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_cli_token_requires_admin(
|
||||
integration_client: AsyncClient,
|
||||
) -> None:
|
||||
"""No admin token / no bearer → 403."""
|
||||
resp = await integration_client.post(
|
||||
"/admin/api/cli-tokens", json={"name": "no-auth"}
|
||||
)
|
||||
assert resp.status_code == 403
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
# Listing
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_cli_tokens_returns_preview_not_full_token(
|
||||
admin_client: AsyncClient,
|
||||
) -> None:
|
||||
"""Listing never leaks the raw token."""
|
||||
create = await admin_client.post(
|
||||
"/admin/api/cli-tokens", json={"name": "secret-keeper"}
|
||||
)
|
||||
assert create.status_code == 200
|
||||
full_token = create.json()["token"]
|
||||
|
||||
resp = await admin_client.get("/admin/api/cli-tokens")
|
||||
assert resp.status_code == 200
|
||||
items = resp.json()
|
||||
assert any(t["name"] == "secret-keeper" for t in items)
|
||||
|
||||
for t in items:
|
||||
# No 'token' field, only 'token_preview'
|
||||
assert "token" not in t
|
||||
assert "token_preview" in t
|
||||
assert full_token not in t["token_preview"]
|
||||
assert "..." in t["token_preview"]
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_cli_tokens_requires_admin(
|
||||
integration_client: AsyncClient,
|
||||
) -> None:
|
||||
resp = await integration_client.get("/admin/api/cli-tokens")
|
||||
assert resp.status_code == 403
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
# Using a CLI token as admin auth
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_cli_token_authorizes_admin_endpoints(
|
||||
admin_client: AsyncClient,
|
||||
integration_client: AsyncClient,
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""A freshly-created CLI token can be used as Bearer on admin endpoints."""
|
||||
create = await admin_client.post(
|
||||
"/admin/api/cli-tokens", json={"name": "cli-auth"}
|
||||
)
|
||||
assert create.status_code == 200
|
||||
cli_token = create.json()["token"]
|
||||
token_id = create.json()["id"]
|
||||
|
||||
# Use a NEW client to isolate the header from admin_session_token
|
||||
integration_client.headers["Authorization"] = f"Bearer {cli_token}"
|
||||
resp = await integration_client.get("/admin/api/cli-tokens")
|
||||
assert resp.status_code == 200
|
||||
|
||||
# last_used_at should be populated after use
|
||||
row = await integration_session.get(CliToken, token_id)
|
||||
assert row is not None
|
||||
assert row.last_used_at is not None
|
||||
assert row.last_used_at >= row.created_at
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_expired_cli_token_is_rejected(
|
||||
admin_client: AsyncClient,
|
||||
integration_client: AsyncClient,
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""A CLI token with expires_at in the past → 403."""
|
||||
create = await admin_client.post(
|
||||
"/admin/api/cli-tokens",
|
||||
json={"name": "will-expire", "expires_in_days": 1},
|
||||
)
|
||||
assert create.status_code == 200
|
||||
cli_token = create.json()["token"]
|
||||
token_id = create.json()["id"]
|
||||
|
||||
# Force-expire it in the DB
|
||||
row = await integration_session.get(CliToken, token_id)
|
||||
assert row is not None
|
||||
row.expires_at = int(time.time()) - 1
|
||||
integration_session.add(row)
|
||||
await integration_session.commit()
|
||||
|
||||
integration_client.headers["Authorization"] = f"Bearer {cli_token}"
|
||||
resp = await integration_client.get("/admin/api/cli-tokens")
|
||||
assert resp.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalid_bearer_token_is_rejected(
|
||||
integration_client: AsyncClient,
|
||||
) -> None:
|
||||
integration_client.headers["Authorization"] = "Bearer not-a-real-token"
|
||||
resp = await integration_client.get("/admin/api/cli-tokens")
|
||||
assert resp.status_code == 403
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
# Revocation
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_revoke_cli_token_removes_auth(
|
||||
admin_client: AsyncClient,
|
||||
integration_client: AsyncClient,
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""After DELETE, the token no longer authorizes."""
|
||||
create = await admin_client.post(
|
||||
"/admin/api/cli-tokens", json={"name": "to-revoke"}
|
||||
)
|
||||
token_id = create.json()["id"]
|
||||
cli_token = create.json()["token"]
|
||||
|
||||
revoke = await admin_client.delete(f"/admin/api/cli-tokens/{token_id}")
|
||||
assert revoke.status_code == 200
|
||||
assert revoke.json() == {"ok": True, "deleted_id": token_id}
|
||||
|
||||
# Row is gone
|
||||
row = await integration_session.get(CliToken, token_id)
|
||||
assert row is None
|
||||
|
||||
# Can no longer be used for auth
|
||||
integration_client.headers["Authorization"] = f"Bearer {cli_token}"
|
||||
resp = await integration_client.get("/admin/api/cli-tokens")
|
||||
assert resp.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_revoke_unknown_cli_token_returns_404(
|
||||
admin_client: AsyncClient,
|
||||
) -> None:
|
||||
resp = await admin_client.delete("/admin/api/cli-tokens/does-not-exist")
|
||||
assert resp.status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_revoke_cli_token_requires_admin(
|
||||
integration_client: AsyncClient,
|
||||
) -> None:
|
||||
resp = await integration_client.delete("/admin/api/cli-tokens/anything")
|
||||
assert resp.status_code == 403
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
# Lifecycle / uniqueness
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_multiple_tokens_are_independent(
|
||||
admin_client: AsyncClient,
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Creating N tokens yields N unique tokens that all live in DB."""
|
||||
names = ["dev-a", "dev-b", "dev-c"]
|
||||
raw_tokens: list[str] = []
|
||||
ids: list[str] = []
|
||||
for name in names:
|
||||
r = await admin_client.post(
|
||||
"/admin/api/cli-tokens", json={"name": name}
|
||||
)
|
||||
assert r.status_code == 200
|
||||
raw_tokens.append(r.json()["token"])
|
||||
ids.append(r.json()["id"])
|
||||
|
||||
# All unique
|
||||
assert len(set(raw_tokens)) == len(raw_tokens)
|
||||
assert len(set(ids)) == len(ids)
|
||||
|
||||
# All in DB
|
||||
result = await integration_session.exec(
|
||||
select(CliToken).where(CliToken.name.in_(names)) # type: ignore[attr-defined]
|
||||
)
|
||||
rows = result.all()
|
||||
assert {r.name for r in rows} == set(names)
|
||||
@@ -267,10 +267,9 @@ async def test_admin_endpoint_unauthenticated(
|
||||
"""Test GET /admin/ endpoint redirects to /"""
|
||||
await db_snapshot.capture()
|
||||
|
||||
response = await integration_client.get("/admin/")
|
||||
response = await integration_client.get("/admin/api/settings")
|
||||
|
||||
assert response.status_code == 307
|
||||
assert response.headers.get("location") == "/"
|
||||
assert response.status_code == 403
|
||||
|
||||
diff = await db_snapshot.diff()
|
||||
assert len(diff["api_keys"]["added"]) == 0
|
||||
|
||||
274
tests/integration/test_insufficient_balance.py
Normal file
274
tests/integration/test_insufficient_balance.py
Normal file
@@ -0,0 +1,274 @@
|
||||
"""
|
||||
Tests showing how a user hits "Insufficient balance: X mSats required for this model"
|
||||
when their balance is too low for the model's cost.
|
||||
|
||||
The log line that triggered this:
|
||||
WARNING Insufficient billing balance during validation
|
||||
ERROR Bearer token validation failed: HTTPException: 402:
|
||||
{'error': {'message': 'Insufficient balance: 622888 mSats required
|
||||
for this model. 20320 available.', ...}}
|
||||
|
||||
This happens in validate_bearer_key (auth.py) when:
|
||||
billing_key.total_balance < min_cost (model's max cost)
|
||||
|
||||
and also in pay_for_request when the atomic UPDATE finds no available balance.
|
||||
"""
|
||||
|
||||
import uuid
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from httpx import AsyncClient
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.db import ApiKey
|
||||
|
||||
|
||||
def _key(balance: int, reserved: int = 0) -> ApiKey:
|
||||
return ApiKey(
|
||||
hashed_key=f"test_{uuid.uuid4().hex}",
|
||||
balance=balance,
|
||||
reserved_balance=reserved,
|
||||
total_spent=0,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 1 — simplest case: balance < model cost → pay_for_request raises 402
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pay_for_request_raises_402_when_balance_too_low(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""
|
||||
User has 20_000 msats. Model costs 622_888 msats.
|
||||
pay_for_request must raise HTTP 402 with a clear message.
|
||||
"""
|
||||
from routstr.auth import pay_for_request
|
||||
|
||||
model_cost = 622_888
|
||||
user_balance = 20_000
|
||||
|
||||
key = _key(balance=user_balance)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await pay_for_request(key, model_cost, integration_session)
|
||||
|
||||
assert exc_info.value.status_code == 402
|
||||
detail = exc_info.value.detail
|
||||
assert isinstance(detail, dict)
|
||||
error = detail["error"]
|
||||
assert error["code"] == "insufficient_balance"
|
||||
assert str(model_cost) in error["message"]
|
||||
assert str(user_balance) in error["message"]
|
||||
|
||||
# Balance must be untouched
|
||||
await integration_session.refresh(key)
|
||||
assert key.balance == user_balance
|
||||
assert key.reserved_balance == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 2 — balance is zero
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pay_for_request_raises_402_on_zero_balance(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""User with zero balance cannot make any request."""
|
||||
from routstr.auth import pay_for_request
|
||||
|
||||
key = _key(balance=0)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await pay_for_request(key, 1_000, integration_session)
|
||||
|
||||
assert exc_info.value.status_code == 402
|
||||
detail = exc_info.value.detail
|
||||
assert isinstance(detail, dict)
|
||||
assert detail["error"]["code"] == "insufficient_balance"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 3 — all balance is reserved (total_balance = balance - reserved = 0)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pay_for_request_raises_402_when_all_balance_reserved(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""
|
||||
User has 50_000 msats balance but 50_000 is already reserved for in-flight
|
||||
requests. Free balance (total_balance) = 0. Should get 402.
|
||||
"""
|
||||
from routstr.auth import pay_for_request
|
||||
|
||||
key = _key(balance=50_000, reserved=50_000)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await pay_for_request(key, 1_000, integration_session)
|
||||
|
||||
assert exc_info.value.status_code == 402
|
||||
# Balance and reserved must be untouched
|
||||
await integration_session.refresh(key)
|
||||
assert key.balance == 50_000
|
||||
assert key.reserved_balance == 50_000
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_balance_info_matches_chat_available_balance(
|
||||
integration_client: AsyncClient,
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""
|
||||
Regression for /v1/balance/info showing gross funds while chat admission
|
||||
rejects with a negative available balance.
|
||||
"""
|
||||
from routstr.auth import pay_for_request
|
||||
|
||||
key = _key(balance=4_404_339, reserved=4_410_636)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
response = await integration_client.get(
|
||||
"/v1/balance/info",
|
||||
headers={"Authorization": f"Bearer sk-{key.hashed_key}"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["balance"] == -6_297
|
||||
assert body["reserved"] == 4_410_636
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await pay_for_request(key, 1, integration_session)
|
||||
|
||||
assert exc_info.value.status_code == 402
|
||||
detail = exc_info.value.detail
|
||||
assert isinstance(detail, dict)
|
||||
assert "-6297 available" in detail["error"]["message"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 4 — balance just one msat below model cost
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pay_for_request_raises_402_one_msat_short(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Off-by-one: balance is exactly model_cost - 1."""
|
||||
from routstr.auth import pay_for_request
|
||||
|
||||
model_cost = 10_000
|
||||
key = _key(balance=model_cost - 1)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await pay_for_request(key, model_cost, integration_session)
|
||||
|
||||
assert exc_info.value.status_code == 402
|
||||
await integration_session.refresh(key)
|
||||
assert key.balance == model_cost - 1 # untouched
|
||||
assert key.reserved_balance == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 5 — balance exactly equal to model cost → succeeds
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pay_for_request_succeeds_when_balance_equals_cost(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Balance == model cost: the request should be reserved successfully."""
|
||||
from routstr.auth import pay_for_request
|
||||
|
||||
model_cost = 10_000
|
||||
key = _key(balance=model_cost)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
# Should not raise
|
||||
await pay_for_request(key, model_cost, integration_session)
|
||||
|
||||
await integration_session.refresh(key)
|
||||
assert key.reserved_balance == model_cost
|
||||
assert key.balance == model_cost # balance unchanged, only reserved goes up
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 6 — HTTP layer returns 402 JSON with the right shape
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_http_402_response_shape_on_insufficient_balance(
|
||||
integration_client: AsyncClient,
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""
|
||||
End-to-end: POST /v1/chat/completions with a key whose balance is far below
|
||||
the mocked model cost returns HTTP 402 with the expected JSON error body.
|
||||
|
||||
Matches exactly the log snippet in the bug report:
|
||||
'Insufficient balance: X mSats required for this model. Y available.'
|
||||
"""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
model_cost = 622_888
|
||||
user_balance = 20_320
|
||||
|
||||
key = _key(balance=user_balance)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
# Minimal model stub so proxy routing doesn't 400 before reaching balance check
|
||||
mock_model = MagicMock()
|
||||
mock_model.sats_pricing = None
|
||||
|
||||
# Upstream stub — never reached because balance check fires first
|
||||
mock_upstream = MagicMock()
|
||||
mock_upstream.prepare_headers = MagicMock(return_value={})
|
||||
|
||||
with (
|
||||
patch("routstr.proxy.get_model_instance", return_value=mock_model),
|
||||
patch("routstr.proxy.get_provider_for_model", return_value=[mock_upstream]),
|
||||
# Patch where it is used (proxy imports it at module level)
|
||||
patch(
|
||||
"routstr.proxy.get_max_cost_for_model",
|
||||
new=AsyncMock(return_value=model_cost),
|
||||
),
|
||||
):
|
||||
response = await integration_client.post(
|
||||
"/v1/chat/completions",
|
||||
headers={"Authorization": f"Bearer sk-{key.hashed_key}"},
|
||||
json={
|
||||
"model": "gpt-4o",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 402
|
||||
body = response.json()
|
||||
# FastAPI wraps HTTPException detail under "detail"
|
||||
error = body["detail"]["error"]
|
||||
assert error["code"] == "insufficient_balance"
|
||||
assert error["type"] == "insufficient_quota"
|
||||
assert str(model_cost) in error["message"]
|
||||
assert str(user_balance) in error["message"]
|
||||
|
||||
# Balance must be completely untouched
|
||||
await integration_session.refresh(key)
|
||||
assert key.balance == user_balance
|
||||
assert key.reserved_balance == 0
|
||||
assert key.total_spent == 0
|
||||
142
tests/integration/test_key_logic.py
Normal file
142
tests/integration/test_key_logic.py
Normal file
@@ -0,0 +1,142 @@
|
||||
import time
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
import pytest
|
||||
from sqlmodel import select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.auth import pay_for_request
|
||||
from routstr.core.db import ApiKey
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_validity_date(integration_session: AsyncSession) -> None:
|
||||
# 1. Create a key that is expired
|
||||
expired_time = int(time.time()) - 3600
|
||||
key = ApiKey(hashed_key="expired_key", balance=1000, validity_date=expired_time)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
# 2. Try to pay for a request - should fail
|
||||
with pytest.raises(Exception) as excinfo:
|
||||
await pay_for_request(key, 100, integration_session)
|
||||
assert "expired" in str(excinfo.value).lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_balance_limit(integration_session: AsyncSession) -> None:
|
||||
# 1. Create a key with a balance limit
|
||||
key = ApiKey(
|
||||
hashed_key="limited_key", balance=10000, balance_limit=500, total_spent=450
|
||||
)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
# 2. Try to pay for a request that exceeds the limit
|
||||
with pytest.raises(Exception) as excinfo:
|
||||
await pay_for_request(key, 100, integration_session)
|
||||
assert "limit exceeded" in str(excinfo.value).lower()
|
||||
|
||||
# 3. Try to pay for a request that fits
|
||||
await pay_for_request(key, 50, integration_session)
|
||||
await integration_session.refresh(key)
|
||||
# Note: total_spent is updated in adjust_payment_for_tokens,
|
||||
# but pay_for_request checks it.
|
||||
# In our current logic, pay_for_request checks (total_spent + cost) > balance_limit.
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_daily_reset_policy(integration_session: AsyncSession) -> None:
|
||||
# 1. Create a key with a daily reset policy and old reset date
|
||||
yesterday = int((datetime.now() - timedelta(days=1)).timestamp())
|
||||
key = ApiKey(
|
||||
hashed_key="daily_reset_key",
|
||||
balance=10000,
|
||||
balance_limit=1000,
|
||||
balance_limit_reset="daily",
|
||||
balance_limit_reset_date=yesterday,
|
||||
total_spent=900,
|
||||
)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
# 2. Pay for a request - should trigger reset first because it's a new day
|
||||
# Request is 200, total_spent is 900. 900+200 > 1000,
|
||||
# but reset should happen making total_spent 0, then 0+200 < 1000.
|
||||
await pay_for_request(key, 200, integration_session)
|
||||
|
||||
await integration_session.refresh(key)
|
||||
assert key.total_spent == 0 # Reset in pay_for_request happens before charging
|
||||
# Wait, the charging logic in pay_for_request increments parent/billing_key's total_requests,
|
||||
# but total_spent is updated in adjust_payment_for_tokens.
|
||||
# However, the reset logic sets total_spent to 0.
|
||||
assert key.balance_limit_reset_date is not None
|
||||
assert key.balance_limit_reset_date > yesterday
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_periodic_key_reset_job(integration_session: AsyncSession) -> None:
|
||||
# 1. Create multiple keys needing reset
|
||||
yesterday = int((datetime.now() - timedelta(days=1)).timestamp())
|
||||
key1 = ApiKey(
|
||||
hashed_key="job_reset_key_1",
|
||||
balance=1000,
|
||||
balance_limit=1000,
|
||||
balance_limit_reset="daily",
|
||||
balance_limit_reset_date=yesterday,
|
||||
total_spent=500,
|
||||
)
|
||||
key2 = ApiKey(
|
||||
hashed_key="job_reset_key_2",
|
||||
balance=1000,
|
||||
balance_limit=1000,
|
||||
balance_limit_reset="daily",
|
||||
balance_limit_reset_date=yesterday,
|
||||
total_spent=800,
|
||||
)
|
||||
integration_session.add(key1)
|
||||
integration_session.add(key2)
|
||||
await integration_session.commit()
|
||||
|
||||
# 2. Run the periodic reset logic manually (mocking the background task loop)
|
||||
# We can't easily run the actual loop because it has a sleep,
|
||||
# but we can test the logic inside.
|
||||
|
||||
# Implementation of periodic_key_reset logic for testing:
|
||||
stmt = select(ApiKey).where(ApiKey.balance_limit_reset != None) # noqa: E711
|
||||
keys = (await integration_session.exec(stmt)).all()
|
||||
now = int(time.time())
|
||||
for k in keys:
|
||||
if k.hashed_key in ["job_reset_key_1", "job_reset_key_2"]:
|
||||
k.total_spent = 0
|
||||
k.balance_limit_reset_date = now
|
||||
integration_session.add(k)
|
||||
await integration_session.commit()
|
||||
|
||||
# 3. Verify resets
|
||||
await integration_session.refresh(key1)
|
||||
await integration_session.refresh(key2)
|
||||
assert key1.total_spent == 0
|
||||
assert key2.total_spent == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_does_not_delete_key(integration_session: AsyncSession) -> None:
|
||||
# This requires mocking the router call or testing the logic in balance.py
|
||||
from routstr.balance import ApiKey
|
||||
|
||||
key = ApiKey(hashed_key="refund_test_key", balance=1000, reserved_balance=100)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
# Logic from refund_wallet_endpoint:
|
||||
key.balance = 0
|
||||
key.reserved_balance = 0
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
|
||||
# Verify key still exists
|
||||
fetched_key = await integration_session.get(ApiKey, "refund_test_key")
|
||||
assert fetched_key is not None
|
||||
assert fetched_key.balance == 0
|
||||
assert fetched_key.reserved_balance == 0
|
||||
198
tests/integration/test_lightning_invoice_rip08.py
Normal file
198
tests/integration/test_lightning_invoice_rip08.py
Normal file
@@ -0,0 +1,198 @@
|
||||
"""RIP-08 lightning invoice endpoint tests.
|
||||
|
||||
Verifies both the spec-compliant path (`POST /lightning/invoice` with
|
||||
`Authorization: Bearer sk-...`) and the legacy path
|
||||
(`POST /v1/balance/lightning/invoice` with `api_key` in body).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from httpx import AsyncClient
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.core.db import ApiKey
|
||||
|
||||
RIP08_PATH = "/lightning/invoice"
|
||||
LEGACY_PATH = "/v1/balance/lightning/invoice"
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def patch_invoice_generation() -> Any:
|
||||
"""Stub out `generate_lightning_invoice` so no mint round-trip is needed."""
|
||||
counter = {"n": 0}
|
||||
|
||||
async def fake_generate(amount_sats: int, description: str) -> tuple[str, str]:
|
||||
counter["n"] += 1
|
||||
return (
|
||||
f"lnbc{amount_sats}n1pfakeinvoice{counter['n']}",
|
||||
f"payment_hash_{counter['n']}",
|
||||
)
|
||||
|
||||
with patch(
|
||||
"routstr.lightning.generate_lightning_invoice",
|
||||
side_effect=fake_generate,
|
||||
) as m:
|
||||
yield m
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def seeded_topup_key(integration_session: AsyncSession) -> str:
|
||||
"""Insert an ApiKey row and return the public `sk-...` form."""
|
||||
hashed = "0" * 64
|
||||
key = ApiKey(
|
||||
hashed_key=hashed,
|
||||
balance=0,
|
||||
refund_currency="sat",
|
||||
refund_mint_url="http://localhost:3338",
|
||||
)
|
||||
integration_session.add(key)
|
||||
await integration_session.commit()
|
||||
return f"sk-{hashed}"
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("path", [RIP08_PATH, LEGACY_PATH])
|
||||
async def test_create_invoice_purpose_create(
|
||||
integration_client: AsyncClient,
|
||||
patch_invoice_generation: Any,
|
||||
path: str,
|
||||
) -> None:
|
||||
"""`purpose=create` works on both paths and requires no auth."""
|
||||
resp = await integration_client.post(
|
||||
path,
|
||||
json={"amount_sats": 1000, "purpose": "create"},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
body = resp.json()
|
||||
assert body["amount_sats"] == 1000
|
||||
assert body["bolt11"].startswith("lnbc")
|
||||
assert body["invoice_id"]
|
||||
assert body["payment_hash"]
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("path", [RIP08_PATH, LEGACY_PATH])
|
||||
async def test_topup_with_authorization_header(
|
||||
integration_client: AsyncClient,
|
||||
patch_invoice_generation: Any,
|
||||
seeded_topup_key: str,
|
||||
path: str,
|
||||
) -> None:
|
||||
"""RIP-08: topup using `Authorization: Bearer sk-...` header (no api_key in body)."""
|
||||
resp = await integration_client.post(
|
||||
path,
|
||||
json={"amount_sats": 500, "purpose": "topup"},
|
||||
headers={"Authorization": f"Bearer {seeded_topup_key}"},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
body = resp.json()
|
||||
assert body["amount_sats"] == 500
|
||||
assert body["bolt11"].startswith("lnbc")
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("path", [RIP08_PATH, LEGACY_PATH])
|
||||
async def test_topup_with_legacy_api_key_in_body(
|
||||
integration_client: AsyncClient,
|
||||
patch_invoice_generation: Any,
|
||||
seeded_topup_key: str,
|
||||
path: str,
|
||||
) -> None:
|
||||
"""Legacy: topup with `api_key` in body still accepted on both paths."""
|
||||
resp = await integration_client.post(
|
||||
path,
|
||||
json={
|
||||
"amount_sats": 250,
|
||||
"purpose": "topup",
|
||||
"api_key": seeded_topup_key,
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
assert resp.json()["amount_sats"] == 250
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("path", [RIP08_PATH, LEGACY_PATH])
|
||||
async def test_topup_missing_auth_returns_401(
|
||||
integration_client: AsyncClient,
|
||||
patch_invoice_generation: Any,
|
||||
path: str,
|
||||
) -> None:
|
||||
"""Topup without any credential is rejected on both paths."""
|
||||
resp = await integration_client.post(
|
||||
path,
|
||||
json={"amount_sats": 100, "purpose": "topup"},
|
||||
)
|
||||
assert resp.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("path", [RIP08_PATH, LEGACY_PATH])
|
||||
async def test_topup_unknown_api_key_returns_404(
|
||||
integration_client: AsyncClient,
|
||||
patch_invoice_generation: Any,
|
||||
path: str,
|
||||
) -> None:
|
||||
resp = await integration_client.post(
|
||||
path,
|
||||
json={"amount_sats": 100, "purpose": "topup"},
|
||||
headers={"Authorization": "Bearer sk-deadbeef"},
|
||||
)
|
||||
assert resp.status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("path", [RIP08_PATH, LEGACY_PATH])
|
||||
async def test_invoice_status_404_for_unknown_id(
|
||||
integration_client: AsyncClient,
|
||||
path: str,
|
||||
) -> None:
|
||||
base = path.rsplit("/invoice", 1)[0] + "/invoice"
|
||||
resp = await integration_client.get(f"{base}/does-not-exist/status")
|
||||
assert resp.status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_purpose_defaults_to_create(
|
||||
integration_client: AsyncClient,
|
||||
patch_invoice_generation: Any,
|
||||
) -> None:
|
||||
"""Per RIP-08, `purpose` may be omitted and defaults to `create`."""
|
||||
resp = await integration_client.post(
|
||||
RIP08_PATH,
|
||||
json={"amount_sats": 100},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
assert resp.json()["amount_sats"] == 100
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_authorization_header_overrides_body_api_key(
|
||||
integration_client: AsyncClient,
|
||||
patch_invoice_generation: Any,
|
||||
seeded_topup_key: str,
|
||||
) -> None:
|
||||
"""Header api_key wins over body api_key: bogus body must not cause 404."""
|
||||
resp = await integration_client.post(
|
||||
RIP08_PATH,
|
||||
json={
|
||||
"amount_sats": 100,
|
||||
"purpose": "topup",
|
||||
"api_key": "sk-" + "f" * 64, # bogus body key
|
||||
},
|
||||
headers={"Authorization": f"Bearer {seeded_topup_key}"},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
143
tests/integration/test_provider_fee_enforcement.py
Normal file
143
tests/integration/test_provider_fee_enforcement.py
Normal file
@@ -0,0 +1,143 @@
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any, AsyncGenerator, cast
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from routstr.core.db import AsyncSession, ModelRow, UpstreamProviderRow
|
||||
from routstr.payment.models import Architecture, Model, Pricing
|
||||
from routstr.proxy import refresh_model_maps
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_enforce_lowest_provider_fee_for_same_url(
|
||||
integration_session: Any,
|
||||
) -> None:
|
||||
"""Test that the algorithm selects the provider with the lowest fee when URLs match."""
|
||||
|
||||
# 1. Create two providers with the same URL but different fees
|
||||
url = "https://api.example.com"
|
||||
p1 = UpstreamProviderRow(
|
||||
provider_type="custom",
|
||||
base_url=url,
|
||||
api_key="key1",
|
||||
enabled=True,
|
||||
provider_fee=1.01,
|
||||
)
|
||||
p2 = UpstreamProviderRow(
|
||||
provider_type="custom",
|
||||
base_url=url,
|
||||
api_key="key2",
|
||||
enabled=True,
|
||||
provider_fee=1.05,
|
||||
)
|
||||
|
||||
integration_session.add(p1)
|
||||
integration_session.add(p2)
|
||||
await integration_session.commit()
|
||||
await integration_session.refresh(p1)
|
||||
await integration_session.refresh(p2)
|
||||
|
||||
assert p1.id is not None
|
||||
assert p2.id is not None
|
||||
|
||||
# 2. Add a model for each provider
|
||||
m1 = ModelRow(
|
||||
id="model-a",
|
||||
name="Model A",
|
||||
created=1,
|
||||
description="desc",
|
||||
context_length=100,
|
||||
architecture='{"modality": "text", "input_modalities": ["text"], "output_modalities": ["text"], "tokenizer": "tiktoken", "instruct_type": "chat"}',
|
||||
pricing='{"prompt": 1.0, "completion": 1.0}',
|
||||
upstream_provider_id=p1.id,
|
||||
enabled=True,
|
||||
)
|
||||
m2 = ModelRow(
|
||||
id="model-a",
|
||||
name="Model A",
|
||||
created=1,
|
||||
description="desc",
|
||||
context_length=100,
|
||||
architecture='{"modality": "text", "input_modalities": ["text"], "output_modalities": ["text"], "tokenizer": "tiktoken", "instruct_type": "chat"}',
|
||||
pricing='{"prompt": 1.0, "completion": 1.0}',
|
||||
upstream_provider_id=p2.id,
|
||||
enabled=True,
|
||||
)
|
||||
|
||||
integration_session.add(m1)
|
||||
integration_session.add(m2)
|
||||
await integration_session.commit()
|
||||
|
||||
# 3. Create mock provider instances
|
||||
class MockProvider(BaseUpstreamProvider):
|
||||
db_id: int
|
||||
|
||||
def __init__(self, db_id: int, base_url: str, api_key: str, fee: float):
|
||||
super().__init__(base_url, api_key, fee)
|
||||
self.db_id = db_id
|
||||
self.provider_type = "custom"
|
||||
|
||||
def get_cached_models(self) -> list[Model]:
|
||||
return [
|
||||
Model(
|
||||
id="model-a",
|
||||
name="Model A",
|
||||
created=1,
|
||||
description="desc",
|
||||
context_length=100,
|
||||
architecture=Architecture(
|
||||
modality="text",
|
||||
input_modalities=["text"],
|
||||
output_modalities=["text"],
|
||||
tokenizer="tiktoken",
|
||||
instruct_type="chat",
|
||||
),
|
||||
pricing=Pricing(prompt=1.0, completion=1.0),
|
||||
enabled=True,
|
||||
upstream_provider_id=self.db_id,
|
||||
)
|
||||
]
|
||||
|
||||
async def refresh_models_cache(self) -> None:
|
||||
pass
|
||||
|
||||
def prepare_headers(self, request_headers: dict[str, str]) -> dict[str, str]:
|
||||
return request_headers
|
||||
|
||||
# 4. Inject mock providers into the proxy
|
||||
from routstr import proxy
|
||||
|
||||
assert p1.id is not None
|
||||
assert p2.id is not None
|
||||
|
||||
# Need to patch proxy._upstreams and proxy.create_session
|
||||
mp1: MockProvider = MockProvider(p1.id, url, "key1", 1.01)
|
||||
mp2: MockProvider = MockProvider(p2.id, url, "key2", 1.05)
|
||||
|
||||
with (
|
||||
patch("routstr.proxy._upstreams", [mp1, mp2]),
|
||||
patch("routstr.proxy.create_session") as mock_session_factory,
|
||||
):
|
||||
# Configure mock_session_factory to return a session that uses the test engine
|
||||
@asynccontextmanager
|
||||
async def mock_create_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
yield integration_session
|
||||
|
||||
mock_session_factory.return_value = mock_create_session()
|
||||
|
||||
await refresh_model_maps()
|
||||
|
||||
# 5. Check which provider is selected for 'model-a'
|
||||
provider_map = proxy.get_provider_for_model("model-a")
|
||||
|
||||
# Assertions
|
||||
assert provider_map is not None
|
||||
assert len(provider_map) >= 1
|
||||
|
||||
# Check the first one, cast to MockProvider to access db_id
|
||||
best_provider = cast(MockProvider, provider_map[0])
|
||||
assert best_provider.db_id == p1.id
|
||||
assert best_provider.provider_fee == 1.01
|
||||
@@ -3,13 +3,17 @@ Integration tests for provider management functionality.
|
||||
Tests GET /v1/providers/ endpoint for listing and managing providers.
|
||||
"""
|
||||
|
||||
import time
|
||||
from types import TracebackType
|
||||
from typing import Any, Generator
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from httpx import AsyncClient
|
||||
|
||||
from routstr.discovery import _PROVIDERS_CACHE
|
||||
from routstr.core.admin import admin_sessions
|
||||
from routstr.core.db import UpstreamProviderRow
|
||||
from routstr.nostr.discovery import _PROVIDERS_CACHE
|
||||
|
||||
from .utils import ResponseValidator
|
||||
|
||||
@@ -71,9 +75,10 @@ async def test_providers_endpoint_default_response(
|
||||
}
|
||||
|
||||
with patch(
|
||||
"routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events
|
||||
"routstr.nostr.discovery.query_nostr_relay_for_providers",
|
||||
return_value=mock_events,
|
||||
):
|
||||
with patch("routstr.discovery.fetch_provider_health") as mock_fetch:
|
||||
with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch:
|
||||
# Configure mock to return appropriate responses
|
||||
mock_fetch.side_effect = lambda url: mock_fetch_responses.get(
|
||||
url, {"status_code": 500, "json": {"error": "Unknown provider"}}
|
||||
@@ -135,9 +140,10 @@ async def test_providers_endpoint_with_include_json(
|
||||
}
|
||||
|
||||
with patch(
|
||||
"routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events
|
||||
"routstr.nostr.discovery.query_nostr_relay_for_providers",
|
||||
return_value=mock_events,
|
||||
):
|
||||
with patch("routstr.discovery.fetch_provider_health") as mock_fetch:
|
||||
with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch:
|
||||
mock_fetch.return_value = {
|
||||
"status_code": 200,
|
||||
"json": mock_provider_response,
|
||||
@@ -209,9 +215,10 @@ async def test_providers_data_structure_validation(
|
||||
}
|
||||
|
||||
with patch(
|
||||
"routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events
|
||||
"routstr.nostr.discovery.query_nostr_relay_for_providers",
|
||||
return_value=mock_events,
|
||||
):
|
||||
with patch("routstr.discovery.fetch_provider_health") as mock_fetch:
|
||||
with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch:
|
||||
mock_fetch.return_value = mock_health_response
|
||||
|
||||
response = await integration_client.get("/v1/providers/?include_json=true")
|
||||
@@ -256,7 +263,8 @@ async def test_providers_endpoint_no_providers_found(
|
||||
]
|
||||
|
||||
with patch(
|
||||
"routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events
|
||||
"routstr.nostr.discovery.query_nostr_relay_for_providers",
|
||||
return_value=mock_events,
|
||||
):
|
||||
response = await integration_client.get("/v1/providers/")
|
||||
|
||||
@@ -317,10 +325,11 @@ async def test_providers_endpoint_offline_providers(
|
||||
}
|
||||
|
||||
with patch(
|
||||
"routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events
|
||||
"routstr.nostr.discovery.query_nostr_relay_for_providers",
|
||||
return_value=mock_events,
|
||||
):
|
||||
with patch(
|
||||
"routstr.discovery.fetch_provider_health",
|
||||
"routstr.nostr.discovery.fetch_provider_health",
|
||||
side_effect=mock_fetch_provider_health,
|
||||
):
|
||||
response = await integration_client.get("/v1/providers/?include_json=true")
|
||||
@@ -386,9 +395,10 @@ async def test_providers_endpoint_duplicate_urls(
|
||||
]
|
||||
|
||||
with patch(
|
||||
"routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events
|
||||
"routstr.nostr.discovery.query_nostr_relay_for_providers",
|
||||
return_value=mock_events,
|
||||
):
|
||||
with patch("routstr.discovery.fetch_provider_health") as mock_fetch:
|
||||
with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch:
|
||||
mock_fetch.return_value = {
|
||||
"status_code": 200,
|
||||
"endpoint": "root",
|
||||
@@ -425,7 +435,8 @@ async def test_providers_endpoint_nostr_relay_failures(
|
||||
raise Exception("Connection to relay failed")
|
||||
|
||||
with patch(
|
||||
"routstr.discovery.query_nostr_relay_for_providers", side_effect=failing_query
|
||||
"routstr.nostr.discovery.query_nostr_relay_for_providers",
|
||||
side_effect=failing_query,
|
||||
):
|
||||
response = await integration_client.get("/v1/providers/")
|
||||
|
||||
@@ -463,9 +474,10 @@ async def test_providers_endpoint_malformed_urls(
|
||||
]
|
||||
|
||||
with patch(
|
||||
"routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events
|
||||
"routstr.nostr.discovery.query_nostr_relay_for_providers",
|
||||
return_value=mock_events,
|
||||
):
|
||||
with patch("routstr.discovery.fetch_provider_health") as mock_fetch:
|
||||
with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch:
|
||||
mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}}
|
||||
|
||||
response = await integration_client.get("/v1/providers/")
|
||||
@@ -495,9 +507,10 @@ async def test_providers_endpoint_response_format(
|
||||
]
|
||||
|
||||
with patch(
|
||||
"routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events
|
||||
"routstr.nostr.discovery.query_nostr_relay_for_providers",
|
||||
return_value=mock_events,
|
||||
):
|
||||
with patch("routstr.discovery.fetch_provider_health") as mock_fetch:
|
||||
with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch:
|
||||
mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}}
|
||||
|
||||
# Test default format
|
||||
@@ -545,9 +558,10 @@ async def test_providers_endpoint_concurrent_requests(
|
||||
]
|
||||
|
||||
with patch(
|
||||
"routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events
|
||||
"routstr.nostr.discovery.query_nostr_relay_for_providers",
|
||||
return_value=mock_events,
|
||||
):
|
||||
with patch("routstr.discovery.fetch_provider_health") as mock_fetch:
|
||||
with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch:
|
||||
mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}}
|
||||
|
||||
# Create concurrent requests
|
||||
@@ -587,9 +601,10 @@ async def test_providers_endpoint_parameter_validation(
|
||||
]
|
||||
|
||||
with patch(
|
||||
"routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events
|
||||
"routstr.nostr.discovery.query_nostr_relay_for_providers",
|
||||
return_value=mock_events,
|
||||
):
|
||||
with patch("routstr.discovery.fetch_provider_health") as mock_fetch:
|
||||
with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch:
|
||||
mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}}
|
||||
|
||||
# Test various parameter values
|
||||
@@ -639,9 +654,10 @@ async def test_no_database_changes_during_provider_operations(
|
||||
]
|
||||
|
||||
with patch(
|
||||
"routstr.discovery.query_nostr_relay_for_providers", return_value=mock_events
|
||||
"routstr.nostr.discovery.query_nostr_relay_for_providers",
|
||||
return_value=mock_events,
|
||||
):
|
||||
with patch("routstr.discovery.fetch_provider_health") as mock_fetch:
|
||||
with patch("routstr.nostr.discovery.fetch_provider_health") as mock_fetch:
|
||||
mock_fetch.return_value = {"status_code": 200, "json": {"status": "online"}}
|
||||
|
||||
# Make multiple requests with different parameters
|
||||
@@ -666,3 +682,88 @@ async def test_no_database_changes_during_provider_operations(
|
||||
assert final_diff["api_keys"]["added"] == []
|
||||
assert final_diff["api_keys"]["modified"] == []
|
||||
assert final_diff["api_keys"]["removed"] == []
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_admin_routstr_topup_retries_transient_upstream_failure(
|
||||
integration_client: AsyncClient,
|
||||
integration_session: Any,
|
||||
) -> None:
|
||||
admin_token = "test-admin-token"
|
||||
admin_sessions[admin_token] = int(time.time()) + 3600
|
||||
integration_client.headers["Authorization"] = f"Bearer {admin_token}"
|
||||
|
||||
provider = UpstreamProviderRow(
|
||||
provider_type="routstr",
|
||||
base_url="https://node.example",
|
||||
api_key="sk-upstream-test",
|
||||
enabled=True,
|
||||
provider_fee=1.01,
|
||||
)
|
||||
integration_session.add(provider)
|
||||
await integration_session.commit()
|
||||
await integration_session.refresh(provider)
|
||||
|
||||
class MockResponse:
|
||||
def __init__(self, status_code: int, data: dict[str, Any] | None = None):
|
||||
self.status_code = status_code
|
||||
self._data = data or {}
|
||||
self.text = str(self._data)
|
||||
|
||||
def json(self) -> dict[str, Any]:
|
||||
return self._data
|
||||
|
||||
class MockAsyncClient:
|
||||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
|
||||
async def __aenter__(self) -> "MockAsyncClient":
|
||||
return self
|
||||
|
||||
async def __aexit__(
|
||||
self,
|
||||
exc_type: type[BaseException] | None,
|
||||
exc: BaseException | None,
|
||||
tb: TracebackType | None,
|
||||
) -> None:
|
||||
return None
|
||||
|
||||
async def post(
|
||||
self, url: str, json: dict[str, Any], headers: dict[str, str]
|
||||
) -> MockResponse:
|
||||
self.calls += 1
|
||||
assert url == "https://node.example/v1/balance/lightning/invoice"
|
||||
assert json["amount_sats"] == 10
|
||||
assert json["purpose"] == "topup"
|
||||
assert json["api_key"] == "sk-upstream-test"
|
||||
assert headers["Authorization"] == "Bearer sk-upstream-test"
|
||||
|
||||
if self.calls == 1:
|
||||
return MockResponse(500, {"detail": "warmup failure"})
|
||||
|
||||
return MockResponse(
|
||||
200,
|
||||
{
|
||||
"bolt11": "lnbc1testinvoice",
|
||||
"invoice_id": "invoice-123",
|
||||
},
|
||||
)
|
||||
|
||||
mock_client = MockAsyncClient()
|
||||
|
||||
try:
|
||||
with patch("httpx.AsyncClient", return_value=mock_client):
|
||||
response = await integration_client.post(
|
||||
f"/admin/api/upstream-providers/{provider.id}/topup",
|
||||
json={"amount": 10},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
assert data["ok"] is True
|
||||
assert data["topup_data"]["payment_request"] == "lnbc1testinvoice"
|
||||
assert data["topup_data"]["invoice_id"] == "invoice-123"
|
||||
assert mock_client.calls == 2
|
||||
finally:
|
||||
admin_sessions.pop(admin_token, None)
|
||||
|
||||
@@ -171,10 +171,11 @@ async def test_proxy_get_unauthorized_access(integration_client: AsyncClient) ->
|
||||
assert response.status_code == 200 # GET requests are allowed
|
||||
|
||||
# Test 2: POST requests without auth should return 401
|
||||
# Note: Model validation happens before auth, so missing model returns 400
|
||||
response = await integration_client.post(
|
||||
"/v1/chat/completions", json={"test": "data"}
|
||||
)
|
||||
assert response.status_code == 401
|
||||
assert response.status_code in [400, 401] # Accept both for now
|
||||
|
||||
# Test 3: POST with invalid API key
|
||||
# Note: After refactor, model validation may happen before auth validation
|
||||
@@ -550,9 +551,6 @@ async def test_proxy_get_concurrent_requests(
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_proxy_get_response_format_preservation(
|
||||
|
||||
268
tests/integration/test_reservation_lifecycle.py
Normal file
268
tests/integration/test_reservation_lifecycle.py
Normal file
@@ -0,0 +1,268 @@
|
||||
"""
|
||||
Tests for the reservation lifecycle:
|
||||
|
||||
1. Reserve → reserved_balance increases, available (total_balance) decreases.
|
||||
2. Reserve → revert → reserved_balance restored, balance untouched.
|
||||
3. Reserve → finalise → reserved_balance released, balance charged.
|
||||
4. Two parallel reserves, only one fits → second blocked with 402.
|
||||
5. Three parallel reserves, two fit, third blocked with 402.
|
||||
6. Sequential reserves until balance exhausted → next request blocked.
|
||||
|
||||
Reservation invariant enforced by the atomic WHERE clause in pay_for_request:
|
||||
balance - reserved_balance >= cost_per_request
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import uuid
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from routstr.auth import pay_for_request, revert_pay_for_request
|
||||
from routstr.core.db import ApiKey, create_session
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _make_key(balance: int, reserved: int = 0) -> ApiKey:
|
||||
return ApiKey(
|
||||
hashed_key=f"test_{uuid.uuid4().hex}",
|
||||
balance=balance,
|
||||
reserved_balance=reserved,
|
||||
total_spent=0,
|
||||
total_requests=0,
|
||||
)
|
||||
|
||||
|
||||
async def _persist(session: AsyncSession, key: ApiKey) -> ApiKey:
|
||||
session.add(key)
|
||||
await session.commit()
|
||||
await session.refresh(key)
|
||||
return key
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 1 — Reserve: reserved_balance increases, available balance decreases
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reserve_increases_reserved_balance(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""pay_for_request must increment reserved_balance by cost_per_request."""
|
||||
cost = 100
|
||||
key = await _persist(integration_session, _make_key(balance=500))
|
||||
|
||||
await pay_for_request(key, cost, integration_session)
|
||||
await integration_session.refresh(key)
|
||||
|
||||
assert key.reserved_balance == cost
|
||||
assert key.balance == 500 # balance column is NOT decremented on reserve
|
||||
assert key.total_balance == 500 - cost # available = balance - reserved
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 2 — Revert: reserved_balance restored, balance untouched
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_revert_releases_reservation(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""revert_pay_for_request must release the reservation without touching balance."""
|
||||
cost = 150
|
||||
key = await _persist(integration_session, _make_key(balance=300))
|
||||
|
||||
await pay_for_request(key, cost, integration_session)
|
||||
await integration_session.refresh(key)
|
||||
assert key.reserved_balance == cost
|
||||
|
||||
await revert_pay_for_request(key, integration_session, cost)
|
||||
await integration_session.refresh(key)
|
||||
|
||||
assert key.reserved_balance == 0
|
||||
assert key.balance == 300 # balance unchanged after revert
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 3 — Finalise: reservation released + balance charged
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_finalise_releases_reservation_and_charges_balance(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""adjust_payment_for_tokens must zero reserved_balance and deduct actual cost."""
|
||||
from unittest.mock import patch
|
||||
|
||||
from routstr.auth import adjust_payment_for_tokens
|
||||
from routstr.payment.cost_calculation import CostData
|
||||
|
||||
cost = 100
|
||||
actual = 80 # actual < reserved → refund path
|
||||
key = await _persist(integration_session, _make_key(balance=500))
|
||||
|
||||
await pay_for_request(key, cost, integration_session)
|
||||
await integration_session.refresh(key)
|
||||
assert key.reserved_balance == cost
|
||||
|
||||
cost_data = CostData(
|
||||
base_msats=0,
|
||||
input_msats=40,
|
||||
output_msats=40,
|
||||
total_msats=actual,
|
||||
total_usd=0.0,
|
||||
input_tokens=50,
|
||||
output_tokens=50,
|
||||
)
|
||||
response_data = {"model": "test-model", "usage": {"prompt_tokens": 50, "completion_tokens": 50}}
|
||||
|
||||
with patch("routstr.auth.calculate_cost", return_value=cost_data):
|
||||
await adjust_payment_for_tokens(key, response_data, integration_session, cost)
|
||||
|
||||
await integration_session.refresh(key)
|
||||
|
||||
assert key.reserved_balance == 0
|
||||
assert key.balance == 500 - actual
|
||||
assert key.total_spent == actual
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 4 — Concurrent: second parallel reserve blocked when balance exhausted
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_second_reserve_blocked_when_balance_exhausted(
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
"""When two requests race for the same balance, only one succeeds; the other gets 402."""
|
||||
cost = 300
|
||||
key_hash = f"test_concurrent_{uuid.uuid4().hex}"
|
||||
|
||||
async with create_session() as session:
|
||||
key = ApiKey(
|
||||
hashed_key=key_hash,
|
||||
balance=300, # exactly enough for ONE reservation
|
||||
reserved_balance=0,
|
||||
total_spent=0,
|
||||
total_requests=0,
|
||||
)
|
||||
session.add(key)
|
||||
await session.commit()
|
||||
|
||||
results: list[str] = []
|
||||
|
||||
async def attempt_reserve() -> None:
|
||||
async with create_session() as session:
|
||||
fresh_key = await session.get(ApiKey, key_hash)
|
||||
assert fresh_key is not None
|
||||
try:
|
||||
await pay_for_request(fresh_key, cost, session)
|
||||
results.append("success")
|
||||
except HTTPException as exc:
|
||||
assert exc.status_code == 402
|
||||
results.append("blocked")
|
||||
|
||||
await asyncio.gather(attempt_reserve(), attempt_reserve())
|
||||
|
||||
assert sorted(results) == ["blocked", "success"], (
|
||||
f"Expected exactly one success and one 402, got: {results}"
|
||||
)
|
||||
|
||||
async with create_session() as session:
|
||||
final = await session.get(ApiKey, key_hash)
|
||||
assert final is not None
|
||||
|
||||
# reserved_balance must equal exactly one reservation (not two)
|
||||
assert final.reserved_balance == cost, (
|
||||
f"Expected reserved_balance={cost}, got {final.reserved_balance}"
|
||||
)
|
||||
assert final.balance == 300, "Balance column must not be modified by reservation"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 5 — Concurrent: three requests, two fit, third blocked
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_three_parallel_reserves_third_blocked(
|
||||
patched_db_engine: None,
|
||||
) -> None:
|
||||
"""Balance covers two reservations exactly; the third concurrent request must be blocked."""
|
||||
cost = 100
|
||||
key_hash = f"test_three_parallel_{uuid.uuid4().hex}"
|
||||
|
||||
async with create_session() as session:
|
||||
key = ApiKey(
|
||||
hashed_key=key_hash,
|
||||
balance=200, # fits exactly 2 reservations of 100
|
||||
reserved_balance=0,
|
||||
total_spent=0,
|
||||
total_requests=0,
|
||||
)
|
||||
session.add(key)
|
||||
await session.commit()
|
||||
|
||||
results: list[str] = []
|
||||
|
||||
async def attempt_reserve() -> None:
|
||||
async with create_session() as session:
|
||||
fresh_key = await session.get(ApiKey, key_hash)
|
||||
assert fresh_key is not None
|
||||
try:
|
||||
await pay_for_request(fresh_key, cost, session)
|
||||
results.append("success")
|
||||
except HTTPException as exc:
|
||||
assert exc.status_code == 402
|
||||
results.append("blocked")
|
||||
|
||||
await asyncio.gather(
|
||||
attempt_reserve(),
|
||||
attempt_reserve(),
|
||||
attempt_reserve(),
|
||||
)
|
||||
|
||||
successes = results.count("success")
|
||||
blocked = results.count("blocked")
|
||||
|
||||
assert successes == 2, f"Expected 2 successes, got {successes}: {results}"
|
||||
assert blocked == 1, f"Expected 1 blocked, got {blocked}: {results}"
|
||||
|
||||
async with create_session() as session:
|
||||
final = await session.get(ApiKey, key_hash)
|
||||
assert final is not None
|
||||
|
||||
assert final.reserved_balance == cost * 2, (
|
||||
f"Expected reserved_balance={cost * 2}, got {final.reserved_balance}"
|
||||
)
|
||||
assert final.balance == 200, "Balance column must not be modified by reservation"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test 6 — Sequential exhaustion: reserve until empty, next request blocked
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sequential_reserves_block_when_balance_exhausted(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Repeated reservations should block as soon as available balance drops below cost."""
|
||||
cost = 100
|
||||
key = await _persist(integration_session, _make_key(balance=250))
|
||||
|
||||
# First two succeed (100 + 100 = 200 ≤ 250)
|
||||
await pay_for_request(key, cost, integration_session)
|
||||
await pay_for_request(key, cost, integration_session)
|
||||
await integration_session.refresh(key)
|
||||
assert key.reserved_balance == 200
|
||||
assert key.total_balance == 50 # 250 - 200
|
||||
|
||||
# Third: only 50 available, need 100 → blocked
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await pay_for_request(key, cost, integration_session)
|
||||
|
||||
assert exc_info.value.status_code == 402
|
||||
await integration_session.refresh(key)
|
||||
assert key.reserved_balance == 200 # unchanged after failed reserve
|
||||
@@ -133,13 +133,16 @@ async def test_reserved_balance_with_successful_requests(
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_insufficient_reserved_balance_for_revert(
|
||||
async def test_revert_with_zero_reserved_balance_is_noop(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Test revert_pay_for_request behavior with insufficient reserved balance."""
|
||||
"""Test that revert_pay_for_request is a no-op when reserved_balance is 0.
|
||||
|
||||
Previously this would drive reserved_balance negative. With the floor guard,
|
||||
it should return False and leave reserved_balance at 0.
|
||||
"""
|
||||
from routstr.auth import revert_pay_for_request
|
||||
|
||||
# Create key with zero reserved balance
|
||||
unique_key = f"test_revert_key_{uuid.uuid4().hex[:8]}"
|
||||
test_key = ApiKey(
|
||||
hashed_key=unique_key,
|
||||
@@ -149,17 +152,211 @@ async def test_insufficient_reserved_balance_for_revert(
|
||||
integration_session.add(test_key)
|
||||
await integration_session.commit()
|
||||
|
||||
# Try to revert more than available
|
||||
# Note: Current implementation allows reserved_balance to go negative
|
||||
await revert_pay_for_request(test_key, integration_session, 100)
|
||||
# Try to revert more than available — should be a no-op
|
||||
result = await revert_pay_for_request(test_key, integration_session, 100)
|
||||
|
||||
# Refresh to get updated values
|
||||
await integration_session.refresh(test_key)
|
||||
|
||||
# Current implementation allows negative reserved balance
|
||||
assert test_key.reserved_balance == -100, (
|
||||
f"Expected reserved_balance to be -100, got: {test_key.reserved_balance}"
|
||||
assert result is False, "Revert should return False when reservation already released"
|
||||
assert test_key.reserved_balance == 0, (
|
||||
f"Reserved balance should remain 0, got: {test_key.reserved_balance}"
|
||||
)
|
||||
assert test_key.total_requests == -1, (
|
||||
f"Expected total_requests to be -1, got: {test_key.total_requests}"
|
||||
assert test_key.total_requests == 0, (
|
||||
f"Total requests should remain 0, got: {test_key.total_requests}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_revert_with_sufficient_reserved_balance_succeeds(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Test that revert_pay_for_request works correctly when there is enough reserved balance."""
|
||||
from routstr.auth import revert_pay_for_request
|
||||
|
||||
unique_key = f"test_revert_ok_{uuid.uuid4().hex[:8]}"
|
||||
test_key = ApiKey(
|
||||
hashed_key=unique_key,
|
||||
balance=5000,
|
||||
reserved_balance=500,
|
||||
total_requests=3,
|
||||
)
|
||||
integration_session.add(test_key)
|
||||
await integration_session.commit()
|
||||
|
||||
result = await revert_pay_for_request(test_key, integration_session, 500)
|
||||
|
||||
await integration_session.refresh(test_key)
|
||||
|
||||
assert result is True, "Revert should return True on success"
|
||||
assert test_key.reserved_balance == 0, (
|
||||
f"Reserved balance should be 0, got: {test_key.reserved_balance}"
|
||||
)
|
||||
assert test_key.total_requests == 2, (
|
||||
f"Total requests should be 2, got: {test_key.total_requests}"
|
||||
)
|
||||
assert test_key.balance == 5000, "Balance should not change on revert"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_revert_partial_reserved_balance_is_noop(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Test that reverting more than the current reserved_balance is a no-op."""
|
||||
from routstr.auth import revert_pay_for_request
|
||||
|
||||
unique_key = f"test_revert_partial_{uuid.uuid4().hex[:8]}"
|
||||
test_key = ApiKey(
|
||||
hashed_key=unique_key,
|
||||
balance=5000,
|
||||
reserved_balance=50,
|
||||
total_requests=1,
|
||||
)
|
||||
integration_session.add(test_key)
|
||||
await integration_session.commit()
|
||||
|
||||
# Try to revert 500 when only 50 is reserved — should be no-op
|
||||
result = await revert_pay_for_request(test_key, integration_session, 500)
|
||||
|
||||
await integration_session.refresh(test_key)
|
||||
|
||||
assert result is False, "Revert should fail when cost > reserved_balance"
|
||||
assert test_key.reserved_balance == 50, (
|
||||
f"Reserved balance should stay at 50, got: {test_key.reserved_balance}"
|
||||
)
|
||||
assert test_key.total_requests == 1, (
|
||||
f"Total requests should stay at 1, got: {test_key.total_requests}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_double_revert_prevented(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Test that calling revert twice doesn't drive reserved_balance negative.
|
||||
|
||||
This simulates the double-revert scenario where both upstream/base.py
|
||||
and proxy.py attempt to revert the same reservation.
|
||||
"""
|
||||
from routstr.auth import revert_pay_for_request
|
||||
|
||||
unique_key = f"test_double_revert_{uuid.uuid4().hex[:8]}"
|
||||
test_key = ApiKey(
|
||||
hashed_key=unique_key,
|
||||
balance=10000,
|
||||
reserved_balance=500,
|
||||
total_requests=5,
|
||||
)
|
||||
integration_session.add(test_key)
|
||||
await integration_session.commit()
|
||||
|
||||
# First revert — should succeed
|
||||
result1 = await revert_pay_for_request(test_key, integration_session, 500)
|
||||
await integration_session.refresh(test_key)
|
||||
|
||||
assert result1 is True
|
||||
assert test_key.reserved_balance == 0
|
||||
assert test_key.total_requests == 4
|
||||
|
||||
# Second revert of the same amount — should be no-op
|
||||
result2 = await revert_pay_for_request(test_key, integration_session, 500)
|
||||
await integration_session.refresh(test_key)
|
||||
|
||||
assert result2 is False, "Second revert should be a no-op"
|
||||
assert test_key.reserved_balance == 0, (
|
||||
f"Reserved balance should stay 0, got: {test_key.reserved_balance}"
|
||||
)
|
||||
assert test_key.total_requests == 4, (
|
||||
f"Total requests should stay 4, got: {test_key.total_requests}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sequential_reverts_never_go_negative(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Test that multiple reverts don't cause negative reserved_balance.
|
||||
|
||||
Simulates the double-revert scenario where multiple code paths
|
||||
attempt to revert the same reservation.
|
||||
"""
|
||||
from routstr.auth import revert_pay_for_request
|
||||
|
||||
unique_key = f"test_multi_revert_{uuid.uuid4().hex[:8]}"
|
||||
test_key = ApiKey(
|
||||
hashed_key=unique_key,
|
||||
balance=10000,
|
||||
reserved_balance=500,
|
||||
total_requests=5,
|
||||
)
|
||||
integration_session.add(test_key)
|
||||
await integration_session.commit()
|
||||
|
||||
# Run 5 sequential reverts for the same 500 reservation
|
||||
results = []
|
||||
for _ in range(5):
|
||||
r = await revert_pay_for_request(test_key, integration_session, 500)
|
||||
results.append(r)
|
||||
|
||||
await integration_session.refresh(test_key)
|
||||
|
||||
# Exactly one should succeed, rest should be no-ops
|
||||
success_count = sum(1 for r in results if r is True)
|
||||
assert success_count == 1, (
|
||||
f"Exactly one revert should succeed, got {success_count} successes"
|
||||
)
|
||||
assert test_key.reserved_balance == 0, (
|
||||
f"Reserved balance should be 0, got: {test_key.reserved_balance}"
|
||||
)
|
||||
assert test_key.reserved_balance >= 0, (
|
||||
f"Reserved balance went negative: {test_key.reserved_balance}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_child_key_revert_floor_guard(
|
||||
integration_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Test that child key reserved_balance also has floor guard on revert."""
|
||||
from routstr.auth import revert_pay_for_request
|
||||
|
||||
parent_key_hash = f"test_parent_{uuid.uuid4().hex[:8]}"
|
||||
child_key_hash = f"test_child_{uuid.uuid4().hex[:8]}"
|
||||
|
||||
parent_key = ApiKey(
|
||||
hashed_key=parent_key_hash,
|
||||
balance=10000,
|
||||
reserved_balance=500,
|
||||
total_requests=3,
|
||||
)
|
||||
child_key = ApiKey(
|
||||
hashed_key=child_key_hash,
|
||||
balance=0,
|
||||
reserved_balance=500,
|
||||
total_requests=3,
|
||||
parent_key_hash=parent_key_hash,
|
||||
)
|
||||
integration_session.add(parent_key)
|
||||
integration_session.add(child_key)
|
||||
await integration_session.commit()
|
||||
|
||||
# First revert succeeds
|
||||
result1 = await revert_pay_for_request(child_key, integration_session, 500)
|
||||
await integration_session.refresh(parent_key)
|
||||
await integration_session.refresh(child_key)
|
||||
|
||||
assert result1 is True
|
||||
assert parent_key.reserved_balance == 0
|
||||
assert child_key.reserved_balance == 0
|
||||
|
||||
# Second revert is a no-op for both parent and child
|
||||
result2 = await revert_pay_for_request(child_key, integration_session, 500)
|
||||
await integration_session.refresh(parent_key)
|
||||
await integration_session.refresh(child_key)
|
||||
|
||||
assert result2 is False
|
||||
assert parent_key.reserved_balance == 0, (
|
||||
f"Parent reserved_balance should stay 0, got: {parent_key.reserved_balance}"
|
||||
)
|
||||
assert child_key.reserved_balance == 0, (
|
||||
f"Child reserved_balance should stay 0, got: {child_key.reserved_balance}"
|
||||
)
|
||||
|
||||
@@ -13,7 +13,7 @@ import pytest
|
||||
from httpx import AsyncClient
|
||||
from sqlmodel import select
|
||||
|
||||
from routstr.core.db import ApiKey
|
||||
from routstr.core.db import ApiKey, CashuTransaction
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@@ -65,9 +65,10 @@ async def test_full_balance_refund_returns_cashu_token(
|
||||
except Exception as e:
|
||||
pytest.fail(f"Invalid Cashu token format: {e}")
|
||||
|
||||
# Try to use the API key - should fail since it's been deleted
|
||||
# Try to use the API key - should still work but have 0 balance
|
||||
response = await authenticated_client.get("/v1/wallet/")
|
||||
assert response.status_code == 401
|
||||
assert response.status_code == 200
|
||||
assert response.json()["balance"] == 0
|
||||
|
||||
# The refund token has been validated above by decoding it
|
||||
# The API key deletion has been verified by the 401 response
|
||||
@@ -261,17 +262,16 @@ async def test_database_state_after_refund(
|
||||
response = await authenticated_client.post("/v1/wallet/refund")
|
||||
assert response.status_code == 200
|
||||
|
||||
# Verify key is deleted after refund
|
||||
result = await integration_session.execute(
|
||||
select(ApiKey).where(ApiKey.hashed_key == hashed_key) # type: ignore[arg-type]
|
||||
)
|
||||
assert result.scalar_one_or_none() is None
|
||||
# Refresh the key to get the updated balance from the database
|
||||
await integration_session.refresh(key_before)
|
||||
|
||||
# Count total keys to ensure only the specific one was deleted
|
||||
# Verify key balance is 0 after refund
|
||||
assert key_before.balance == 0
|
||||
|
||||
# Count total keys to ensure it wasn't deleted
|
||||
result = await integration_session.execute(select(ApiKey))
|
||||
remaining_keys = result.scalars().all()
|
||||
# Should have no keys left (assuming clean test environment)
|
||||
assert len(remaining_keys) == 0
|
||||
assert len(remaining_keys) == 1
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@@ -356,6 +356,70 @@ async def test_concurrent_refund_requests(
|
||||
assert len(successful) + len(failed) == 5
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_rejects_concurrent_topup_on_same_key(
|
||||
authenticated_client: AsyncClient,
|
||||
testmint_wallet: Any,
|
||||
) -> None:
|
||||
"""Test refund returns 409 when a concurrent topup changes the balance first."""
|
||||
from routstr import balance as balance_module
|
||||
|
||||
wallet_response = await authenticated_client.get("/v1/wallet/")
|
||||
assert wallet_response.status_code == 200
|
||||
initial_balance = wallet_response.json()["balance"]
|
||||
|
||||
topup_amount_sat = 500
|
||||
topup_token = await testmint_wallet.mint_tokens(topup_amount_sat)
|
||||
|
||||
key_looked_up = asyncio.Event()
|
||||
allow_refund_to_continue = asyncio.Event()
|
||||
original_lookup = balance_module._lookup_key_no_create
|
||||
delayed_once = False
|
||||
|
||||
async def delayed_lookup_key_no_create(*args: Any, **kwargs: Any) -> ApiKey | None:
|
||||
nonlocal delayed_once
|
||||
key = await original_lookup(*args, **kwargs)
|
||||
if not delayed_once and key is not None:
|
||||
delayed_once = True
|
||||
key_looked_up.set()
|
||||
await allow_refund_to_continue.wait()
|
||||
return key
|
||||
|
||||
async def issue_refund() -> Any:
|
||||
return await authenticated_client.post("/v1/wallet/refund")
|
||||
|
||||
async def issue_topup() -> Any:
|
||||
await key_looked_up.wait()
|
||||
try:
|
||||
return await authenticated_client.post(
|
||||
"/v1/wallet/topup", params={"cashu_token": topup_token}
|
||||
)
|
||||
finally:
|
||||
allow_refund_to_continue.set()
|
||||
|
||||
with patch(
|
||||
"routstr.balance._lookup_key_no_create", new=delayed_lookup_key_no_create
|
||||
):
|
||||
refund_response, topup_response = await asyncio.gather(
|
||||
issue_refund(), issue_topup()
|
||||
)
|
||||
|
||||
assert topup_response.status_code == 200
|
||||
assert topup_response.json()["msats"] == topup_amount_sat * 1000
|
||||
assert refund_response.status_code == 409
|
||||
assert (
|
||||
refund_response.json()["detail"]
|
||||
== "Balance changed concurrently. Please retry the refund."
|
||||
)
|
||||
|
||||
final_balance_response = await authenticated_client.get("/v1/wallet/")
|
||||
assert final_balance_response.status_code == 200
|
||||
assert final_balance_response.json()["balance"] == (
|
||||
initial_balance + topup_amount_sat * 1000
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_during_active_usage(
|
||||
@@ -388,9 +452,46 @@ async def test_refund_during_active_usage(
|
||||
# Refund should succeed
|
||||
assert refund_response.status_code == 200
|
||||
|
||||
# Further usage should fail
|
||||
# Further usage should return 200 but with 0 balance
|
||||
response = await authenticated_client.get("/v1/wallet/")
|
||||
assert response.status_code == 401
|
||||
assert response.status_code == 200
|
||||
assert response.json()["balance"] == 0
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
async def test_wallet_history_returns_apikey_transactions(
|
||||
authenticated_client: AsyncClient,
|
||||
testmint_wallet: Any,
|
||||
integration_session: Any,
|
||||
) -> None:
|
||||
wallet_response = await authenticated_client.get("/v1/wallet/")
|
||||
api_key = wallet_response.json()["api_key"]
|
||||
hashed_key = api_key[3:] if api_key.startswith("sk-") else api_key
|
||||
|
||||
topup_token = await testmint_wallet.mint_tokens(250)
|
||||
topup_response = await authenticated_client.post(
|
||||
"/v1/wallet/topup", params={"cashu_token": topup_token}
|
||||
)
|
||||
assert topup_response.status_code == 200
|
||||
|
||||
refund_response = await authenticated_client.post("/v1/wallet/refund")
|
||||
assert refund_response.status_code == 200
|
||||
|
||||
history_response = await authenticated_client.get("/v1/wallet/history")
|
||||
assert history_response.status_code == 200
|
||||
transactions = history_response.json()["transactions"]
|
||||
assert len(transactions) >= 2
|
||||
assert all("api_key_hashed_key" not in tx for tx in transactions)
|
||||
assert {tx["type"] for tx in transactions} >= {"in", "out"}
|
||||
|
||||
db_result = await integration_session.execute(
|
||||
select(CashuTransaction).where(
|
||||
CashuTransaction.api_key_hashed_key == hashed_key
|
||||
)
|
||||
)
|
||||
db_transactions = db_result.scalars().all()
|
||||
assert len(db_transactions) >= 2
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@@ -535,6 +636,3 @@ async def test_refund_with_expired_key(
|
||||
# Should still allow manual refund
|
||||
assert response.status_code == 200
|
||||
assert response.json()["recipient"] == "expired@ln.address"
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -11,7 +11,7 @@ import pytest
|
||||
from httpx import AsyncClient
|
||||
from sqlmodel import select
|
||||
|
||||
from routstr.core.db import ApiKey
|
||||
from routstr.core.db import ApiKey, CashuTransaction
|
||||
|
||||
from .utils import (
|
||||
CashuTokenGenerator,
|
||||
@@ -71,6 +71,16 @@ async def test_topup_with_valid_token( # type: ignore[no-untyped-def]
|
||||
assert db_key.balance == new_balance
|
||||
assert db_key.balance == initial_balance + (topup_amount * 1000)
|
||||
|
||||
tx_result = await integration_session.execute(
|
||||
select(CashuTransaction).where(
|
||||
CashuTransaction.token == token,
|
||||
CashuTransaction.type == "in",
|
||||
)
|
||||
)
|
||||
tx = tx_result.scalar_one()
|
||||
assert tx.api_key_hashed_key == hashed_key
|
||||
assert tx.source == "apikey"
|
||||
|
||||
|
||||
@pytest.mark.integration
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -1,16 +1,19 @@
|
||||
"""Tests for the model prioritization algorithm."""
|
||||
|
||||
import os
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
|
||||
# Set required env vars before importing
|
||||
os.environ["UPSTREAM_BASE_URL"] = "http://test"
|
||||
os.environ["UPSTREAM_API_KEY"] = "test"
|
||||
|
||||
from routstr.algorithm import ( # noqa: E402
|
||||
calculate_model_cost_score,
|
||||
create_model_mappings,
|
||||
get_provider_penalty,
|
||||
should_prefer_model,
|
||||
)
|
||||
from routstr.payment.models import Architecture, Model, Pricing # noqa: E402
|
||||
|
||||
@@ -46,11 +49,21 @@ def create_test_model(
|
||||
)
|
||||
|
||||
|
||||
def create_test_provider(name: str, base_url: str = "http://test.com") -> Mock:
|
||||
def create_test_provider(
|
||||
name: str,
|
||||
base_url: str = "http://test.com",
|
||||
*,
|
||||
db_id: int | None = None,
|
||||
models: list[Model] | None = None,
|
||||
upstream_name: str | None = None,
|
||||
) -> Mock:
|
||||
"""Helper to create a test provider mock."""
|
||||
provider = Mock()
|
||||
provider.provider_type = name
|
||||
provider.base_url = base_url
|
||||
provider.db_id = db_id
|
||||
provider.upstream_name = upstream_name or name
|
||||
provider.get_cached_models.return_value = models or []
|
||||
return provider
|
||||
|
||||
|
||||
@@ -102,98 +115,78 @@ def test_get_provider_penalty_openrouter() -> None:
|
||||
assert penalty == 1.001
|
||||
|
||||
|
||||
def test_should_prefer_model_cheaper_wins() -> None:
|
||||
"""Test that cheaper model is preferred."""
|
||||
cheap_model = create_test_model("cheap", prompt_price=0.001, completion_price=0.002)
|
||||
expensive_model = create_test_model(
|
||||
"expensive", prompt_price=0.03, completion_price=0.06
|
||||
def test_create_model_mappings_includes_db_override_for_missing_cached_model(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Model overrides should still map when provider discovery misses the model."""
|
||||
provider = create_test_provider(
|
||||
"azure",
|
||||
"https://example.openai.azure.com/openai/v1",
|
||||
db_id=7,
|
||||
models=[],
|
||||
)
|
||||
override_model = create_test_model("azure/gpt-4o")
|
||||
override_model.canonical_slug = "azure-deployment"
|
||||
|
||||
def fake_row_to_model(*args, **kwargs) -> Model: # type: ignore[no-untyped-def]
|
||||
return override_model
|
||||
|
||||
monkeypatch.setattr("routstr.payment.models._row_to_model", fake_row_to_model)
|
||||
|
||||
override_row = SimpleNamespace(id="azure/gpt-4o", upstream_provider_id=7, enabled=True)
|
||||
|
||||
model_instances, provider_map, unique_models = create_model_mappings(
|
||||
upstreams=[provider],
|
||||
overrides_by_id={"azure/gpt-4o": (override_row, 1.01)},
|
||||
disabled_model_ids=set(),
|
||||
)
|
||||
|
||||
provider1 = create_test_provider("provider1")
|
||||
provider2 = create_test_provider("provider2")
|
||||
assert "azure/gpt-4o" in model_instances
|
||||
assert provider_map["azure/gpt-4o"] == [provider]
|
||||
assert "gpt-4o" in unique_models
|
||||
|
||||
# Cheaper model should win
|
||||
assert should_prefer_model(
|
||||
cheap_model, provider1, expensive_model, provider2, "test-alias"
|
||||
|
||||
def test_create_model_mappings_dedupes_with_provider_identity_not_provider_type(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Different provider instances of same type should both survive dedupe."""
|
||||
provider_a_model = create_test_model(
|
||||
"azure/gpt-4o", prompt_price=0.01, completion_price=0.01
|
||||
)
|
||||
provider_a = create_test_provider(
|
||||
"azure",
|
||||
"https://a.openai.azure.com/openai/v1",
|
||||
db_id=1,
|
||||
models=[provider_a_model],
|
||||
upstream_name="azure-a",
|
||||
)
|
||||
provider_b = create_test_provider(
|
||||
"azure",
|
||||
"https://b.openai.azure.com/openai/v1",
|
||||
db_id=2,
|
||||
models=[],
|
||||
upstream_name="azure-b",
|
||||
)
|
||||
|
||||
# More expensive model should not win
|
||||
assert not should_prefer_model(
|
||||
expensive_model, provider2, cheap_model, provider1, "test-alias"
|
||||
override_model = create_test_model(
|
||||
"azure/gpt-4o", prompt_price=0.001, completion_price=0.001
|
||||
)
|
||||
override_model.canonical_slug = "azure-b-deployment"
|
||||
|
||||
def fake_row_to_model(*args, **kwargs) -> Model: # type: ignore[no-untyped-def]
|
||||
return override_model
|
||||
|
||||
monkeypatch.setattr("routstr.payment.models._row_to_model", fake_row_to_model)
|
||||
|
||||
override_row = SimpleNamespace(id="azure/gpt-4o", upstream_provider_id=2, enabled=True)
|
||||
|
||||
_, provider_map, _ = create_model_mappings(
|
||||
upstreams=[provider_a, provider_b],
|
||||
overrides_by_id={"azure/gpt-4o": (override_row, 1.01)},
|
||||
disabled_model_ids=set(),
|
||||
)
|
||||
|
||||
|
||||
def test_should_prefer_model_exact_match_wins() -> None:
|
||||
"""Test that exact alias match beats cheaper price."""
|
||||
# Make model IDs match the alias differently
|
||||
exact_match = create_test_model(
|
||||
"test-model", prompt_price=0.03, completion_price=0.06
|
||||
)
|
||||
no_match = create_test_model(
|
||||
"other-model", prompt_price=0.001, completion_price=0.002
|
||||
)
|
||||
|
||||
provider1 = create_test_provider("provider1")
|
||||
provider2 = create_test_provider("provider2")
|
||||
|
||||
# Exact match should win even though it's more expensive
|
||||
assert should_prefer_model(
|
||||
exact_match, provider1, no_match, provider2, "test-model"
|
||||
)
|
||||
|
||||
|
||||
def test_should_prefer_model_openrouter_slight_penalty() -> None:
|
||||
"""Test that OpenRouter has slight penalty compared to other providers."""
|
||||
model1 = create_test_model("model1", prompt_price=0.001, completion_price=0.002)
|
||||
model2 = create_test_model("model2", prompt_price=0.001, completion_price=0.002)
|
||||
|
||||
regular_provider = create_test_provider("regular", "http://provider.com")
|
||||
openrouter_provider = create_test_provider(
|
||||
"openrouter", "https://openrouter.ai/api/v1"
|
||||
)
|
||||
|
||||
# Regular provider should be preferred over OpenRouter at same cost
|
||||
assert should_prefer_model(
|
||||
model1, regular_provider, model2, openrouter_provider, "test-alias"
|
||||
)
|
||||
|
||||
# OpenRouter should not replace regular provider at same cost
|
||||
assert not should_prefer_model(
|
||||
model2, openrouter_provider, model1, regular_provider, "test-alias"
|
||||
)
|
||||
|
||||
|
||||
def test_should_prefer_model_openrouter_can_win_if_cheaper() -> None:
|
||||
"""Test that OpenRouter can still win if significantly cheaper."""
|
||||
cheap_model = create_test_model(
|
||||
"cheap", prompt_price=0.0001, completion_price=0.0002
|
||||
)
|
||||
expensive_model = create_test_model(
|
||||
"expensive", prompt_price=0.03, completion_price=0.06
|
||||
)
|
||||
|
||||
regular_provider = create_test_provider("regular", "http://provider.com")
|
||||
openrouter_provider = create_test_provider(
|
||||
"openrouter", "https://openrouter.ai/api/v1"
|
||||
)
|
||||
|
||||
# OpenRouter should win if it's much cheaper (even with penalty)
|
||||
assert should_prefer_model(
|
||||
cheap_model,
|
||||
openrouter_provider,
|
||||
expensive_model,
|
||||
regular_provider,
|
||||
"test-alias",
|
||||
)
|
||||
|
||||
|
||||
def test_should_prefer_model_same_cost_first_wins() -> None:
|
||||
"""Test that when costs are identical, current model is kept."""
|
||||
model1 = create_test_model("model1", prompt_price=0.001, completion_price=0.002)
|
||||
model2 = create_test_model("model2", prompt_price=0.001, completion_price=0.002)
|
||||
|
||||
provider1 = create_test_provider("provider1")
|
||||
provider2 = create_test_provider("provider2")
|
||||
|
||||
# When costs are equal, should not replace
|
||||
assert not should_prefer_model(model2, provider2, model1, provider1, "test-alias")
|
||||
providers_for_alias = provider_map["azure/gpt-4o"]
|
||||
assert provider_a in providers_for_alias
|
||||
assert provider_b in providers_for_alias
|
||||
assert len(providers_for_alias) == 2
|
||||
|
||||
403
tests/unit/test_balance.py
Normal file
403
tests/unit/test_balance.py
Normal file
@@ -0,0 +1,403 @@
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from routstr.balance import refund_wallet_endpoint
|
||||
from routstr.core.db import ApiKey, CashuTransaction
|
||||
from routstr.wallet import credit_balance
|
||||
|
||||
|
||||
def _make_cashu_tx(
|
||||
token: str,
|
||||
amount: int,
|
||||
unit: str,
|
||||
type: str = "out",
|
||||
request_id: str | None = "req-abc",
|
||||
swept: bool = False,
|
||||
collected: bool = False,
|
||||
) -> CashuTransaction:
|
||||
tx = CashuTransaction(token=token, amount=amount, unit=unit, type=type, request_id=request_id)
|
||||
tx.swept = swept
|
||||
tx.collected = collected
|
||||
return tx
|
||||
|
||||
|
||||
def _exec_result(tx: CashuTransaction | None) -> MagicMock:
|
||||
result = MagicMock()
|
||||
result.first.return_value = tx
|
||||
return result
|
||||
|
||||
|
||||
def _update_result(rowcount: int) -> MagicMock:
|
||||
result = MagicMock()
|
||||
result.rowcount = rowcount
|
||||
return result
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_x_cashu_returns_token() -> None:
|
||||
x_cashu_token = "cashuAtest_token_value"
|
||||
in_tx = _make_cashu_tx(token=x_cashu_token, amount=0, unit="msat", type="in", request_id="req-abc")
|
||||
out_tx = _make_cashu_tx(token="cashuArefund_token", amount=1000, unit="msat", type="out", request_id="req-abc")
|
||||
|
||||
session = MagicMock()
|
||||
session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)])
|
||||
session.add = MagicMock()
|
||||
session.commit = AsyncMock()
|
||||
|
||||
result = await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-somekey",
|
||||
x_cashu=x_cashu_token,
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert isinstance(result, JSONResponse)
|
||||
body = json.loads(result.body)
|
||||
assert body["token"] == "cashuArefund_token"
|
||||
assert body["msats"] == "1000"
|
||||
assert result.headers["X-Cashu"] == "cashuArefund_token"
|
||||
assert out_tx.collected is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_x_cashu_sat_unit() -> None:
|
||||
x_cashu_token = "cashuAsat_token"
|
||||
in_tx = _make_cashu_tx(token=x_cashu_token, amount=0, unit="sat", type="in", request_id="req-sat")
|
||||
out_tx = _make_cashu_tx(token="cashuArefund_sat", amount=500, unit="sat", type="out", request_id="req-sat")
|
||||
|
||||
session = MagicMock()
|
||||
session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)])
|
||||
session.add = MagicMock()
|
||||
session.commit = AsyncMock()
|
||||
|
||||
result = await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-somekey",
|
||||
x_cashu=x_cashu_token,
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert isinstance(result, JSONResponse)
|
||||
body = json.loads(result.body)
|
||||
assert body["token"] == "cashuArefund_sat"
|
||||
assert body["sats"] == "500"
|
||||
assert "msats" not in body
|
||||
assert result.headers["X-Cashu"] == "cashuArefund_sat"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_x_cashu_not_found_raises_404() -> None:
|
||||
from fastapi import HTTPException
|
||||
|
||||
session = MagicMock()
|
||||
session.exec = AsyncMock(return_value=_exec_result(None))
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-somekey",
|
||||
x_cashu="cashuAmissing_token",
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_x_cashu_swept_raises_410() -> None:
|
||||
from fastapi import HTTPException
|
||||
|
||||
in_tx = _make_cashu_tx(token="cashuAswept_token", amount=0, unit="msat", type="in", request_id="req-swept")
|
||||
out_tx = _make_cashu_tx(token="cashuAswept", amount=100, unit="msat", type="out", request_id="req-swept", swept=True)
|
||||
|
||||
session = MagicMock()
|
||||
session.exec = AsyncMock(side_effect=[_exec_result(in_tx), _exec_result(out_tx)])
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-somekey",
|
||||
x_cashu="cashuAswept_token",
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 410
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# source field defaults
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_cashu_transaction_source_defaults_to_x_cashu() -> None:
|
||||
tx = CashuTransaction(token="cashuAtest", amount=100, unit="msat")
|
||||
assert tx.source == "x-cashu"
|
||||
|
||||
|
||||
def test_cashu_transaction_source_can_be_apikey() -> None:
|
||||
tx = CashuTransaction(token="cashuAtest", amount=100, unit="msat", source="apikey")
|
||||
assert tx.source == "apikey"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# apikey-based refund: token logging and CashuTransaction storage
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_api_key(
|
||||
balance: int = 5000,
|
||||
refund_currency: str | None = "sat",
|
||||
refund_mint_url: str | None = "https://mint.example.com",
|
||||
refund_address: str | None = None,
|
||||
parent_key_hash: str | None = None,
|
||||
) -> ApiKey:
|
||||
key = ApiKey(hashed_key="testhash")
|
||||
key.balance = balance
|
||||
key.reserved_balance = 0
|
||||
key.refund_currency = refund_currency
|
||||
key.refund_mint_url = refund_mint_url
|
||||
key.refund_address = refund_address
|
||||
key.parent_key_hash = parent_key_hash
|
||||
key.total_spent = 0
|
||||
key.total_requests = 0
|
||||
return key
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apikey_refund_stores_cashu_transaction_with_apikey_source() -> None:
|
||||
key = _make_api_key(balance=5000, refund_currency="sat")
|
||||
refund_token = "cashuArefund_apikey_token"
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=key)
|
||||
session.exec = AsyncMock(return_value=_update_result(1))
|
||||
session.add = MagicMock()
|
||||
session.commit = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)),
|
||||
patch("routstr.balance.store_cashu_transaction", AsyncMock()) as mock_store,
|
||||
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
|
||||
patch("routstr.balance._refund_cache_set", AsyncMock()),
|
||||
):
|
||||
result = await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-testhash",
|
||||
x_cashu=None,
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert isinstance(result, dict)
|
||||
assert result["token"] == refund_token
|
||||
|
||||
mock_store.assert_awaited_once()
|
||||
call_kwargs = mock_store.call_args.kwargs
|
||||
assert call_kwargs["source"] == "apikey"
|
||||
assert call_kwargs["token"] == refund_token
|
||||
assert call_kwargs["typ"] == "out"
|
||||
assert call_kwargs["api_key_hashed_key"] == key.hashed_key
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apikey_refund_logs_token() -> None:
|
||||
key = _make_api_key(balance=5000, refund_currency="sat")
|
||||
refund_token = "cashuAlogged_token"
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=key)
|
||||
session.exec = AsyncMock(return_value=_update_result(1))
|
||||
session.add = MagicMock()
|
||||
session.commit = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)),
|
||||
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
|
||||
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
|
||||
patch("routstr.balance._refund_cache_set", AsyncMock()),
|
||||
patch("routstr.balance.logger") as mock_logger,
|
||||
):
|
||||
await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-testhash",
|
||||
x_cashu=None,
|
||||
session=session,
|
||||
)
|
||||
|
||||
calls = [str(c) for c in mock_logger.info.call_args_list]
|
||||
assert any("cashu token issued" in c for c in calls)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apikey_refund_log_includes_path() -> None:
|
||||
key = _make_api_key(balance=5000, refund_currency="sat")
|
||||
refund_token = "cashuApath_token"
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=key)
|
||||
session.exec = AsyncMock(return_value=_update_result(1))
|
||||
session.add = MagicMock()
|
||||
session.commit = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.send_token", AsyncMock(return_value=refund_token)),
|
||||
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
|
||||
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
|
||||
patch("routstr.balance._refund_cache_set", AsyncMock()),
|
||||
patch("routstr.balance.logger") as mock_logger,
|
||||
):
|
||||
await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-testhash",
|
||||
x_cashu=None,
|
||||
session=session,
|
||||
)
|
||||
|
||||
# Find the "cashu token issued" call and verify extra contains the path
|
||||
token_issued_calls = [
|
||||
c for c in mock_logger.info.call_args_list
|
||||
if c.args and "cashu token issued" in c.args[0]
|
||||
]
|
||||
assert len(token_issued_calls) == 1
|
||||
extra = token_issued_calls[0].kwargs.get("extra", {})
|
||||
assert extra.get("path") == "/v1/wallet/refund"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apikey_refund_rejects_on_concurrent_balance_change() -> None:
|
||||
"""When the debit CAS fails (rowcount=0), no token is minted and 409 is returned."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
key = _make_api_key(balance=5000, refund_currency="sat")
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=key)
|
||||
# Debit returns rowcount=0 → balance changed concurrently
|
||||
session.exec = AsyncMock(return_value=_update_result(0))
|
||||
session.commit = AsyncMock()
|
||||
|
||||
mock_send_token = AsyncMock(return_value="cashuAshould_not_be_minted")
|
||||
|
||||
with (
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.send_token", mock_send_token),
|
||||
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
|
||||
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
|
||||
patch("routstr.balance._refund_cache_set", AsyncMock()),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-testhash",
|
||||
x_cashu=None,
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 409
|
||||
# Crucially: send_token must NOT have been called
|
||||
mock_send_token.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_credit_balance_stores_apikey_transaction_history() -> None:
|
||||
key = _make_api_key(balance=1000)
|
||||
session = MagicMock()
|
||||
session.exec = AsyncMock(return_value=_update_result(1))
|
||||
session.commit = AsyncMock()
|
||||
session.refresh = AsyncMock()
|
||||
|
||||
with (
|
||||
patch(
|
||||
"routstr.wallet.recieve_token",
|
||||
AsyncMock(return_value=(100, "sat", "https://mint.example")),
|
||||
),
|
||||
patch("routstr.wallet.store_cashu_transaction", AsyncMock()) as mock_store,
|
||||
):
|
||||
amount = await credit_balance("cashuAtopup_token", key, session)
|
||||
|
||||
assert amount == 100_000
|
||||
mock_store.assert_awaited_once()
|
||||
call_kwargs = mock_store.call_args.kwargs
|
||||
assert call_kwargs["typ"] == "in"
|
||||
assert call_kwargs["source"] == "apikey"
|
||||
assert call_kwargs["api_key_hashed_key"] == key.hashed_key
|
||||
assert call_kwargs["amount"] == 100
|
||||
assert call_kwargs["unit"] == "sat"
|
||||
assert call_kwargs["token"] == "cashuAtopup_token"
|
||||
assert call_kwargs["mint_url"] == "https://mint.example"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apikey_refund_restores_balance_on_mint_failure() -> None:
|
||||
"""When debit succeeds but minting fails, balance must be restored."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
key = _make_api_key(balance=5000, refund_currency="sat")
|
||||
|
||||
# First exec call = debit (succeeds), second = restore
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=key)
|
||||
session.exec = AsyncMock(side_effect=[_update_result(1), _update_result(1)])
|
||||
session.commit = AsyncMock()
|
||||
|
||||
with (
|
||||
patch("routstr.balance.get_billing_key", AsyncMock(return_value=key)),
|
||||
patch("routstr.balance.send_token", AsyncMock(side_effect=Exception("mint down"))),
|
||||
patch("routstr.balance.store_cashu_transaction", AsyncMock()),
|
||||
patch("routstr.balance._refund_cache_get", AsyncMock(return_value=None)),
|
||||
patch("routstr.balance._refund_cache_set", AsyncMock()),
|
||||
patch("routstr.balance.logger"),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-testhash",
|
||||
x_cashu=None,
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 503
|
||||
# Verify two exec calls: debit + restore
|
||||
assert session.exec.await_count == 2
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# no-create guarantee: fresh Cashu/unknown sk- tokens must not create API keys
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_fresh_cashu_bearer_returns_401() -> None:
|
||||
"""Fresh Cashu token not in DB must get 401, never create a new ApiKey."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=None)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await refund_wallet_endpoint(
|
||||
authorization="Bearer cashuAfresh_never_deposited_token",
|
||||
x_cashu=None,
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
session.get.assert_awaited_once()
|
||||
# No add/commit → no key was persisted
|
||||
session.add.assert_not_called()
|
||||
session.commit.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refund_unknown_sk_bearer_returns_401() -> None:
|
||||
"""Unknown sk- key not in DB must get 401."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
session = MagicMock()
|
||||
session.get = AsyncMock(return_value=None)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await refund_wallet_endpoint(
|
||||
authorization="Bearer sk-unknownhash",
|
||||
x_cashu=None,
|
||||
session=session,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
session.get.assert_awaited_once()
|
||||
345
tests/unit/test_cost_calculation_caching.py
Normal file
345
tests/unit/test_cost_calculation_caching.py
Normal file
@@ -0,0 +1,345 @@
|
||||
"""Tests for cache token handling in cost calculation.
|
||||
|
||||
Covers OpenAI vs Anthropic caching formats, edge cases, and billing accuracy.
|
||||
"""
|
||||
|
||||
import os
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
|
||||
os.environ.setdefault("UPSTREAM_API_KEY", "test")
|
||||
os.environ.setdefault("LIGHTNING_ADDRESS", "test@stm.to")
|
||||
|
||||
from routstr.core.settings import settings
|
||||
from routstr.payment.cost_calculation import CostData, MaxCostData, calculate_cost
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_session() -> AsyncMock:
|
||||
"""Mock AsyncSession for cost calculation tests."""
|
||||
return AsyncMock()
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def mock_fixed_pricing(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Mock settings and price to use fixed pricing."""
|
||||
monkeypatch.setattr(settings, "fixed_pricing", True)
|
||||
monkeypatch.setattr(settings, "fixed_per_1k_input_tokens", 0.001)
|
||||
monkeypatch.setattr(settings, "fixed_per_1k_output_tokens", 0.001)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def patch_sats_usd_price() -> None: # type: ignore[misc]
|
||||
"""Patch sats_usd_price to avoid initialization issues."""
|
||||
with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-5):
|
||||
yield
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 1: OpenAI Cache Format
|
||||
# ============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_cache_subtraction(mock_session: AsyncMock) -> None:
|
||||
"""OpenAI includes cached_tokens in prompt_tokens, subtract them."""
|
||||
response = {
|
||||
"model": "gpt-4",
|
||||
"usage": {
|
||||
"prompt_tokens": 2000, # ← Includes 1000 cached
|
||||
"completion_tokens": 100,
|
||||
"prompt_tokens_details": {
|
||||
"cached_tokens": 1000 # ← Extracted separately
|
||||
}
|
||||
}
|
||||
}
|
||||
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
assert result.input_tokens == 1000 # 2000 - 1000
|
||||
assert result.cache_read_input_tokens == 1000
|
||||
assert result.output_tokens == 100
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 2: Anthropic Cache Format
|
||||
# ============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_cache_additive(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
|
||||
"""Anthropic cache tokens are separate (additive) from input_tokens."""
|
||||
response = {
|
||||
"model": "claude-3-5-sonnet",
|
||||
"usage": {
|
||||
"input_tokens": 500, # ← Regular input only
|
||||
"output_tokens": 100,
|
||||
"cache_creation_input_tokens": 1500, # ← Additive, not included above
|
||||
"cache_read_input_tokens": 0,
|
||||
}
|
||||
}
|
||||
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
assert result.input_tokens == 500
|
||||
assert result.cache_creation_input_tokens == 1500
|
||||
assert result.cache_read_input_tokens == 0
|
||||
assert result.output_tokens == 100
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 3: Invalid Cache (Edge Case)
|
||||
# ============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_cache_read_exceeds_prompt_tokens(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
|
||||
"""Handle buggy upstream reporting cached > prompt_tokens."""
|
||||
response = {
|
||||
"model": "gpt-4",
|
||||
"usage": {
|
||||
"prompt_tokens": 100,
|
||||
"completion_tokens": 50,
|
||||
"prompt_tokens_details": {
|
||||
"cached_tokens": 150 # ← Invalid! Greater than prompt
|
||||
}
|
||||
}
|
||||
}
|
||||
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||
|
||||
# Should not go negative
|
||||
assert isinstance(result, CostData)
|
||||
assert result.input_tokens == 0 # max(0, 100 - 150)
|
||||
assert result.cache_read_input_tokens == 150
|
||||
assert result.output_tokens == 50
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 4: Malformed Token Values
|
||||
# ============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_malformed_cache_tokens_coerce_to_zero(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
|
||||
"""Handle non-numeric cache token values."""
|
||||
response = {
|
||||
"model": "gpt-4",
|
||||
"usage": {
|
||||
"prompt_tokens": 100,
|
||||
"completion_tokens": 50,
|
||||
"cache_read_input_tokens": "-50", # ← String, negative
|
||||
"prompt_tokens_details": {
|
||||
"cached_tokens": "invalid" # ← Non-numeric string
|
||||
}
|
||||
}
|
||||
}
|
||||
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||
|
||||
# Both should coerce to 0
|
||||
assert isinstance(result, CostData)
|
||||
assert result.cache_read_input_tokens == 0
|
||||
assert result.input_tokens == 100 # No subtraction if cache_read = 0
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 5: Anthropic Cache Not Subtracted
|
||||
# ============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_cache_not_subtracted(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
|
||||
"""Anthropic cache fields should NOT be subtracted from input_tokens."""
|
||||
response = {
|
||||
"model": "claude-3-5-sonnet",
|
||||
"usage": {
|
||||
"input_tokens": 500,
|
||||
"completion_tokens": 100,
|
||||
"cache_read_input_tokens": 200, # ← Additive, don't subtract
|
||||
}
|
||||
}
|
||||
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||
|
||||
# Anthropic: input_tokens stays as-is
|
||||
assert isinstance(result, CostData)
|
||||
assert result.input_tokens == 500 # NOT 300
|
||||
assert result.cache_read_input_tokens == 200
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 6: Only Cache Read, No Regular Input
|
||||
# ============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_only_cache_read_tokens(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
|
||||
"""Handle response with only cache read tokens."""
|
||||
response = {
|
||||
"model": "gpt-4",
|
||||
"usage": {
|
||||
"prompt_tokens": 0,
|
||||
"completion_tokens": 50,
|
||||
"prompt_tokens_details": {
|
||||
"cached_tokens": 1000
|
||||
}
|
||||
}
|
||||
}
|
||||
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
assert result.input_tokens == 0 # max(0, 0 - 1000)
|
||||
assert result.cache_read_input_tokens == 1000
|
||||
assert result.output_tokens == 50
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 7: Only Cache Creation
|
||||
# ============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_only_cache_creation_tokens(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
|
||||
"""Handle response with only cache creation tokens (Anthropic)."""
|
||||
response = {
|
||||
"model": "claude-3-5-sonnet",
|
||||
"usage": {
|
||||
"input_tokens": 500,
|
||||
"output_tokens": 100,
|
||||
"cache_creation_input_tokens": 2000,
|
||||
"cache_read_input_tokens": 0,
|
||||
}
|
||||
}
|
||||
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
assert result.input_tokens == 500
|
||||
assert result.cache_creation_input_tokens == 2000
|
||||
assert result.cache_read_input_tokens == 0
|
||||
assert result.output_tokens == 100
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 8: Both Cache Read and Creation
|
||||
# ============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_both_cache_read_and_creation(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
|
||||
"""Handle response with both cache read and creation."""
|
||||
response = {
|
||||
"model": "claude-3-5-sonnet",
|
||||
"usage": {
|
||||
"input_tokens": 300,
|
||||
"output_tokens": 100,
|
||||
"cache_creation_input_tokens": 2000,
|
||||
"cache_read_input_tokens": 500,
|
||||
}
|
||||
}
|
||||
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
assert result.input_tokens == 300
|
||||
assert result.cache_creation_input_tokens == 2000
|
||||
assert result.cache_read_input_tokens == 500
|
||||
assert result.output_tokens == 100
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 9: Token Field Fallback
|
||||
# ============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_field_fallback_order(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
|
||||
"""Verify fallback order for token extraction."""
|
||||
# When prompt_tokens is not present, fall back to input_tokens
|
||||
response = {
|
||||
"model": "gpt-4",
|
||||
"usage": {
|
||||
"input_tokens": 250,
|
||||
"completion_tokens": 50,
|
||||
}
|
||||
}
|
||||
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
assert result.input_tokens == 250
|
||||
assert result.output_tokens == 50
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 10: Float Token Values
|
||||
# ============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_float_token_values_coerced_to_int(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
|
||||
"""Handle float token values by converting to int."""
|
||||
response = {
|
||||
"model": "gpt-4",
|
||||
"usage": {
|
||||
"prompt_tokens": 100.7, # Float
|
||||
"completion_tokens": 50.3, # Float
|
||||
"cache_read_input_tokens": 25.9, # Float
|
||||
}
|
||||
}
|
||||
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
assert result.input_tokens == 100 # Floored
|
||||
assert result.output_tokens == 50 # Floored
|
||||
assert result.cache_read_input_tokens == 25 # Floored
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 11: Boolean Cache Tokens
|
||||
# ============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_boolean_cache_tokens_coerced_to_zero(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
|
||||
"""Handle boolean cache token values by coercing to zero."""
|
||||
response = {
|
||||
"model": "gpt-4",
|
||||
"usage": {
|
||||
"prompt_tokens": 100,
|
||||
"completion_tokens": 50,
|
||||
"cache_read_input_tokens": True, # Boolean
|
||||
}
|
||||
}
|
||||
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
assert result.cache_read_input_tokens == 0 # Boolean coerced to 0
|
||||
assert result.input_tokens == 100 # No subtraction
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 12: Zero Cache Tokens
|
||||
# ============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_zero_cache_tokens(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
|
||||
"""Handle explicit zero cache tokens."""
|
||||
response = {
|
||||
"model": "gpt-4",
|
||||
"usage": {
|
||||
"prompt_tokens": 100,
|
||||
"completion_tokens": 50,
|
||||
"prompt_tokens_details": {
|
||||
"cached_tokens": 0
|
||||
}
|
||||
}
|
||||
}
|
||||
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||
|
||||
assert isinstance(result, CostData)
|
||||
assert result.cache_read_input_tokens == 0
|
||||
assert result.input_tokens == 100
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 13: Missing Usage Block
|
||||
# ============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_usage_block(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
|
||||
"""When usage is missing, return MaxCostData with zero tokens."""
|
||||
response = {"model": "gpt-4", "choices": [{"message": {"content": "test"}}]}
|
||||
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||
|
||||
assert isinstance(result, MaxCostData)
|
||||
assert result.input_tokens == 0
|
||||
assert result.cache_read_input_tokens == 0
|
||||
assert result.output_tokens == 0
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 14: Null Usage Block
|
||||
# ============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_null_usage_block(mock_session: AsyncMock, mock_fixed_pricing: None) -> None:
|
||||
"""When usage is null, return MaxCostData with zero tokens."""
|
||||
response = {"model": "gpt-4", "usage": None}
|
||||
result = await calculate_cost(response, max_cost=100000, session=mock_session)
|
||||
|
||||
assert isinstance(result, MaxCostData)
|
||||
assert result.input_tokens == 0
|
||||
assert result.cache_read_input_tokens == 0
|
||||
177
tests/unit/test_count_tokens_local.py
Normal file
177
tests/unit/test_count_tokens_local.py
Normal file
@@ -0,0 +1,177 @@
|
||||
"""Unit tests for the local count_tokens shim.
|
||||
|
||||
The shim runs whenever an upstream that does not support Anthropic's
|
||||
``/v1/messages`` endpoint is asked for a token count. It must always
|
||||
return a 200 JSON ``{"input_tokens": N}`` response and must never raise.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
from unittest.mock import patch
|
||||
|
||||
from routstr.payment.models import Architecture, Model, Pricing
|
||||
from routstr.upstream import count_tokens as count_tokens_module
|
||||
from routstr.upstream.count_tokens import count_tokens_locally
|
||||
|
||||
|
||||
def _make_model(model_id: str = "anthropic/claude-3-5-sonnet") -> Model:
|
||||
pricing = Pricing(prompt=0.000003, completion=0.000015)
|
||||
architecture = Architecture(
|
||||
modality="text",
|
||||
input_modalities=["text"],
|
||||
output_modalities=["text"],
|
||||
tokenizer="cl100k_base",
|
||||
instruct_type=None,
|
||||
)
|
||||
return Model(
|
||||
id=model_id,
|
||||
name=model_id,
|
||||
created=0,
|
||||
description="",
|
||||
context_length=200_000,
|
||||
architecture=architecture,
|
||||
pricing=pricing,
|
||||
)
|
||||
|
||||
|
||||
def _body(payload: dict[str, Any]) -> bytes:
|
||||
return json.dumps(payload).encode()
|
||||
|
||||
|
||||
def _read_payload(response: Any) -> dict[str, Any]:
|
||||
body = response.body if isinstance(response.body, bytes) else bytes(response.body)
|
||||
return json.loads(body.decode())
|
||||
|
||||
|
||||
def test_returns_input_tokens_for_simple_messages() -> None:
|
||||
model = _make_model()
|
||||
request_body = _body(
|
||||
{
|
||||
"model": model.id,
|
||||
"messages": [{"role": "user", "content": "hello world"}],
|
||||
}
|
||||
)
|
||||
|
||||
response = count_tokens_locally(request_body, model)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.media_type == "application/json"
|
||||
payload = _read_payload(response)
|
||||
assert "input_tokens" in payload
|
||||
assert isinstance(payload["input_tokens"], int)
|
||||
assert payload["input_tokens"] >= 0
|
||||
|
||||
|
||||
def test_falls_back_to_estimator_when_litellm_raises() -> None:
|
||||
model = _make_model()
|
||||
request_body = _body(
|
||||
{
|
||||
"model": model.id,
|
||||
"messages": [{"role": "user", "content": "this is a longer message"}],
|
||||
}
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
count_tokens_module,
|
||||
"_count_with_litellm",
|
||||
side_effect=RuntimeError("boom"),
|
||||
):
|
||||
response = count_tokens_locally(request_body, model)
|
||||
|
||||
assert response.status_code == 200
|
||||
payload = _read_payload(response)
|
||||
assert payload["input_tokens"] >= 1
|
||||
|
||||
|
||||
def test_handles_missing_model_object() -> None:
|
||||
request_body = _body(
|
||||
{
|
||||
"model": "anthropic/claude-3-5-sonnet",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
}
|
||||
)
|
||||
|
||||
response = count_tokens_locally(request_body, None)
|
||||
|
||||
assert response.status_code == 200
|
||||
payload = _read_payload(response)
|
||||
assert payload["input_tokens"] >= 0
|
||||
|
||||
|
||||
def test_handles_empty_request_body() -> None:
|
||||
response = count_tokens_locally(b"", _make_model())
|
||||
|
||||
assert response.status_code == 200
|
||||
payload = _read_payload(response)
|
||||
assert payload["input_tokens"] >= 0
|
||||
|
||||
|
||||
def test_handles_malformed_json() -> None:
|
||||
response = count_tokens_locally(b"not-json", _make_model())
|
||||
|
||||
assert response.status_code == 200
|
||||
payload = _read_payload(response)
|
||||
assert payload["input_tokens"] >= 0
|
||||
|
||||
|
||||
def test_includes_system_prompt_in_count() -> None:
|
||||
model = _make_model()
|
||||
short = _body(
|
||||
{
|
||||
"model": model.id,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
}
|
||||
)
|
||||
with_system = _body(
|
||||
{
|
||||
"model": model.id,
|
||||
"system": "You are a helpful assistant with a long preamble " * 10,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
}
|
||||
)
|
||||
|
||||
short_count = _read_payload(count_tokens_locally(short, model))["input_tokens"]
|
||||
long_count = _read_payload(count_tokens_locally(with_system, model))["input_tokens"]
|
||||
|
||||
assert long_count > short_count
|
||||
|
||||
|
||||
def test_supports_anthropic_system_block_list() -> None:
|
||||
model = _make_model()
|
||||
request_body = _body(
|
||||
{
|
||||
"model": model.id,
|
||||
"system": [{"type": "text", "text": "be terse" * 50}],
|
||||
"messages": [{"role": "user", "content": "ok"}],
|
||||
}
|
||||
)
|
||||
|
||||
response = count_tokens_locally(request_body, model)
|
||||
|
||||
payload = _read_payload(response)
|
||||
assert payload["input_tokens"] > 0
|
||||
|
||||
|
||||
def test_uses_forwarded_model_id_when_present() -> None:
|
||||
model = _make_model("anthropic/claude-3-5-sonnet")
|
||||
model.forwarded_model_id = "claude-3-5-sonnet-20241022"
|
||||
request_body = _body(
|
||||
{
|
||||
"model": "ignored",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
}
|
||||
)
|
||||
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
def _capture(model_name: str, body: dict[str, Any]) -> int:
|
||||
captured["model"] = model_name
|
||||
return 7
|
||||
|
||||
with patch.object(count_tokens_module, "_count_with_litellm", side_effect=_capture):
|
||||
response = count_tokens_locally(request_body, model)
|
||||
|
||||
assert captured["model"] == "claude-3-5-sonnet-20241022"
|
||||
assert _read_payload(response)["input_tokens"] == 7
|
||||
@@ -1,155 +0,0 @@
|
||||
"""Unit tests for model row payload conversion.
|
||||
|
||||
This module tests that _model_to_row_payload correctly serializes model data
|
||||
for database storage. Pricing is stored as-is without fee application.
|
||||
Fees are now applied per-provider when reading from the database.
|
||||
|
||||
Key behaviors tested:
|
||||
1. Pricing is stored as-is without fee application
|
||||
2. All model fields are correctly serialized to JSON
|
||||
3. Optional fields are handled correctly (None values)
|
||||
4. Pricing structure is preserved
|
||||
5. Original model objects are not mutated
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
# Set required env vars before importing
|
||||
os.environ["UPSTREAM_BASE_URL"] = "http://test"
|
||||
os.environ["UPSTREAM_API_KEY"] = "test"
|
||||
|
||||
from routstr.payment.models import ( # noqa: E402
|
||||
Architecture,
|
||||
Model,
|
||||
Pricing,
|
||||
_model_to_row_payload,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def base_architecture() -> Architecture:
|
||||
"""Provide standard architecture for test models."""
|
||||
return Architecture(
|
||||
modality="text",
|
||||
input_modalities=["text"],
|
||||
output_modalities=["text"],
|
||||
tokenizer="gpt",
|
||||
instruct_type="chat",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def standard_pricing() -> Pricing:
|
||||
"""Provide standard USD pricing with known values for testing."""
|
||||
return Pricing(
|
||||
prompt=0.001,
|
||||
completion=0.002,
|
||||
request=0.01,
|
||||
image=0.05,
|
||||
web_search=0.03,
|
||||
internal_reasoning=0.015,
|
||||
max_prompt_cost=10.0,
|
||||
max_completion_cost=20.0,
|
||||
max_cost=30.0,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def standard_model(base_architecture: Architecture, standard_pricing: Pricing) -> Model:
|
||||
"""Create a standard test model with known pricing."""
|
||||
return Model(
|
||||
id="test-model-standard",
|
||||
name="Test Model Standard",
|
||||
created=1234567890,
|
||||
description="A standard test model",
|
||||
context_length=8192,
|
||||
architecture=base_architecture,
|
||||
pricing=standard_pricing,
|
||||
)
|
||||
|
||||
|
||||
def test_pricing_stored_without_fees(standard_model: Model) -> None:
|
||||
"""Verify pricing is stored as-is without any fee application."""
|
||||
payload = _model_to_row_payload(standard_model)
|
||||
pricing_str = payload["pricing"]
|
||||
assert isinstance(pricing_str, str)
|
||||
pricing = json.loads(pricing_str)
|
||||
|
||||
assert pricing["prompt"] == pytest.approx(0.001, rel=1e-9)
|
||||
assert pricing["completion"] == pytest.approx(0.002, rel=1e-9)
|
||||
assert pricing["request"] == pytest.approx(0.01, rel=1e-9)
|
||||
assert pricing["image"] == pytest.approx(0.05, rel=1e-9)
|
||||
assert pricing["web_search"] == pytest.approx(0.03, rel=1e-9)
|
||||
assert pricing["internal_reasoning"] == pytest.approx(0.015, rel=1e-9)
|
||||
assert pricing["max_prompt_cost"] == pytest.approx(10.0, rel=1e-9)
|
||||
assert pricing["max_completion_cost"] == pytest.approx(20.0, rel=1e-9)
|
||||
assert pricing["max_cost"] == pytest.approx(30.0, rel=1e-9)
|
||||
|
||||
|
||||
def test_zero_value_pricing_fields(base_architecture: Architecture) -> None:
|
||||
"""Verify that zero-value pricing fields are stored correctly."""
|
||||
zero_pricing = Pricing(
|
||||
prompt=0.0,
|
||||
completion=0.0,
|
||||
request=0.0,
|
||||
image=0.0,
|
||||
web_search=0.0,
|
||||
internal_reasoning=0.0,
|
||||
max_prompt_cost=0.0,
|
||||
max_completion_cost=0.0,
|
||||
max_cost=0.0,
|
||||
)
|
||||
|
||||
model = Model(
|
||||
id="test-model-zero",
|
||||
name="Test Model Zero",
|
||||
created=1234567890,
|
||||
description="A model with zero pricing",
|
||||
context_length=8192,
|
||||
architecture=base_architecture,
|
||||
pricing=zero_pricing,
|
||||
)
|
||||
|
||||
payload = _model_to_row_payload(model)
|
||||
pricing_str = payload["pricing"]
|
||||
assert isinstance(pricing_str, str)
|
||||
pricing = json.loads(pricing_str)
|
||||
|
||||
assert pricing["prompt"] == pytest.approx(0.0, rel=1e-9)
|
||||
assert pricing["completion"] == pytest.approx(0.0, rel=1e-9)
|
||||
assert pricing["request"] == pytest.approx(0.0, rel=1e-9)
|
||||
|
||||
|
||||
def test_payload_structure_unchanged(standard_model: Model) -> None:
|
||||
"""Verify that payload structure matches expectations."""
|
||||
payload = _model_to_row_payload(standard_model)
|
||||
|
||||
assert "id" in payload
|
||||
assert "name" in payload
|
||||
assert "created" in payload
|
||||
assert "description" in payload
|
||||
assert "context_length" in payload
|
||||
assert "architecture" in payload
|
||||
assert "pricing" in payload
|
||||
assert "sats_pricing" in payload
|
||||
assert "per_request_limits" in payload
|
||||
assert "top_provider" in payload
|
||||
assert "enabled" in payload
|
||||
assert "upstream_provider_id" in payload
|
||||
|
||||
assert isinstance(payload["architecture"], str)
|
||||
assert isinstance(payload["pricing"], str)
|
||||
|
||||
|
||||
def test_original_model_not_mutated(standard_model: Model) -> None:
|
||||
"""Verify that the original model object is not mutated."""
|
||||
original_prompt = standard_model.pricing.prompt
|
||||
original_completion = standard_model.pricing.completion
|
||||
|
||||
_model_to_row_payload(standard_model)
|
||||
|
||||
assert standard_model.pricing.prompt == original_prompt
|
||||
assert standard_model.pricing.completion == original_completion
|
||||
93
tests/unit/test_litellm_routing.py
Normal file
93
tests/unit/test_litellm_routing.py
Normal file
@@ -0,0 +1,93 @@
|
||||
"""Tests for `routstr.upstream.litellm_routing.detect_litellm_prefix`."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from routstr.upstream.litellm_routing import detect_litellm_prefix
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"base_url,expected",
|
||||
[
|
||||
# The bug case: custom row pointing at Fireworks must NOT route to openai/.
|
||||
("https://api.fireworks.ai/inference/v1", "fireworks_ai/"),
|
||||
# Other OpenAI-compatible providers commonly plugged into the custom slot.
|
||||
("https://api.groq.com/openai/v1", "groq/"),
|
||||
("https://api.x.ai/v1", "xai/"),
|
||||
("https://api.deepseek.com/v1", "deepseek/"),
|
||||
("https://api.together.xyz/v1", "together_ai/"),
|
||||
("https://api.perplexity.ai", "perplexity/"),
|
||||
("https://openrouter.ai/api/v1", "openrouter/"),
|
||||
("https://api.mistral.ai/v1", "mistral/"),
|
||||
("https://codestral.mistral.ai/v1", "codestral/"),
|
||||
("https://api.cohere.com/v1", "cohere_chat/"),
|
||||
("https://api.cohere.ai/v1", "cohere_chat/"),
|
||||
("https://api.deepinfra.com/v1/openai", "deepinfra/"),
|
||||
("https://api.cerebras.ai/v1", "cerebras/"),
|
||||
("https://api.sambanova.ai/v1", "sambanova/"),
|
||||
("https://api.moonshot.cn/v1", "moonshot/"),
|
||||
("https://api.moonshot.ai/v1", "moonshot/"),
|
||||
("https://api.studio.nebius.com/v1", "nebius/"),
|
||||
("https://api.novita.ai/v3/openai", "novita/"),
|
||||
("https://api.lambda.ai/v1", "lambda_ai/"),
|
||||
("https://api.aimlapi.com/v1", "aiml/"),
|
||||
("https://api.featherless.ai/v1", "featherless_ai/"),
|
||||
("https://integrate.api.nvidia.com/v1", "nvidia_nim/"),
|
||||
("https://inference.baseten.co/v1", "baseten/"),
|
||||
("https://ai-gateway.vercel.sh/v1", "vercel_ai_gateway/"),
|
||||
("https://api.inference.wandb.ai/v1", "wandb/"),
|
||||
("https://api.poe.com/v1", "poe/"),
|
||||
("https://llm.chutes.ai/v1/", "chutes/"),
|
||||
("https://api.v0.dev/v1", "v0/"),
|
||||
("https://api.hyperbolic.xyz/v1", "hyperbolic/"),
|
||||
("https://api.synthetic.new/openai/v1", "synthetic/"),
|
||||
("https://api.stima.tech/v1", "apertis/"),
|
||||
("https://nano-gpt.com/api/v1", "nano-gpt/"),
|
||||
("https://api.friendli.ai/serverless/v1", "friendliai/"),
|
||||
("https://api.galadriel.com/v1", "galadriel/"),
|
||||
("https://api.llama.com/compat/v1", "meta_llama/"),
|
||||
("https://api.minimax.io/v1", "minimax/"),
|
||||
("https://api.minimaxi.com/v1", "minimax/"),
|
||||
("https://platform.publicai.co/v1", "publicai/"),
|
||||
("https://inference.api.nscale.com/v1", "nscale/"),
|
||||
("https://dashscope-intl.aliyuncs.com/compatible-mode/v1", "dashscope/"),
|
||||
("https://api.endpoints.anyscale.com/v1", "anyscale/"),
|
||||
("https://api.ai21.com/studio/v1", "ai21_chat/"),
|
||||
("https://ark.cn-beijing.volces.com/api/v3", "volcengine/"),
|
||||
("https://api.voyageai.com/v1", "voyage/"),
|
||||
("https://api.jina.ai/v1", "jina_ai/"),
|
||||
("https://api.snowflakecomputing.com", "snowflake/"),
|
||||
("https://my-workspace.databricks.com/serving-endpoints", "databricks/"),
|
||||
("https://huggingface.co/api/inference", "huggingface/"),
|
||||
# First-class providers — URL detection still produces the right prefix
|
||||
# so subclasses without an explicit override stay correct.
|
||||
("https://api.openai.com/v1", "openai/"),
|
||||
("https://api.anthropic.com/v1", "anthropic/"),
|
||||
("https://generativelanguage.googleapis.com/v1beta/openai", "gemini/"),
|
||||
("https://us-central1-aiplatform.googleapis.com/v1", "vertex_ai/"),
|
||||
# Azure ordering: must beat api.openai.com.
|
||||
("https://my-resource.openai.azure.com/openai/deployments/foo", "azure/"),
|
||||
# Ollama hints.
|
||||
("http://localhost:11434/v1", "ollama_chat/"),
|
||||
("http://127.0.0.1:11434/v1", "ollama_chat/"),
|
||||
("http://my-ollama-host:11434/v1", "ollama_chat/"),
|
||||
# Casing and trailing slash normalisation.
|
||||
("HTTPS://API.FIREWORKS.AI/INFERENCE/V1/", "fireworks_ai/"),
|
||||
# Unknown host falls back to openai/ (still OpenAI-compatible by convention).
|
||||
("https://example.com/v1", "openai/"),
|
||||
("", "openai/"),
|
||||
],
|
||||
)
|
||||
def test_detect_litellm_prefix(base_url: str, expected: str) -> None:
|
||||
assert detect_litellm_prefix(base_url) == expected
|
||||
|
||||
|
||||
def test_detect_litellm_prefix_none_uses_default() -> None:
|
||||
assert detect_litellm_prefix(None) == "openai/"
|
||||
|
||||
|
||||
def test_detect_litellm_prefix_custom_default() -> None:
|
||||
assert detect_litellm_prefix("https://example.com", default="anthropic/") == (
|
||||
"anthropic/"
|
||||
)
|
||||
333
tests/unit/test_messages_dispatch_cost_accumulation.py
Normal file
333
tests/unit/test_messages_dispatch_cost_accumulation.py
Normal file
@@ -0,0 +1,333 @@
|
||||
"""Tests for cost accumulation in streaming message dispatch.
|
||||
|
||||
Verifies that costs are correctly summed across multiple streaming events,
|
||||
not taking only the maximum.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
|
||||
os.environ.setdefault("UPSTREAM_API_KEY", "test")
|
||||
|
||||
from routstr.upstream.messages_dispatch import annotate_event
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 1: Cost Accumulation via AnnotatedEvent
|
||||
# ============================================================================
|
||||
@pytest.mark.unit
|
||||
def test_annotate_event_extracts_costs() -> None:
|
||||
"""Each event should report its own costs."""
|
||||
event = {
|
||||
"type": "content_block_start",
|
||||
"usage": {"total_cost": 0.010, "input_cost": 0.008, "output_cost": 0.002}
|
||||
}
|
||||
result = annotate_event(event, None)
|
||||
|
||||
assert result.total_cost == 0.010
|
||||
assert result.input_cost == 0.008
|
||||
assert result.output_cost == 0.002
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 2: Multiple Events with Incremental Costs
|
||||
# ============================================================================
|
||||
@pytest.mark.unit
|
||||
def test_multiple_events_sum_costs() -> None:
|
||||
"""When processing multiple events, costs should accumulate (not max)."""
|
||||
# Event 1: initial costs
|
||||
event1 = {
|
||||
"type": "message_start",
|
||||
"usage": {"input_tokens": 100, "total_cost": 0.010}
|
||||
}
|
||||
result1 = annotate_event(event1, None)
|
||||
|
||||
# Event 2: additional costs
|
||||
event2 = {
|
||||
"type": "content_block_start",
|
||||
"usage": {"output_tokens": 50, "total_cost": 0.005}
|
||||
}
|
||||
result2 = annotate_event(event2, None)
|
||||
|
||||
# Event 3: more costs
|
||||
event3 = {
|
||||
"type": "message_delta",
|
||||
"usage": {"output_tokens": 25, "total_cost": 0.008}
|
||||
}
|
||||
result3 = annotate_event(event3, None)
|
||||
|
||||
# Each event should report its own cost (before accumulation)
|
||||
assert result1.total_cost == 0.010
|
||||
assert result2.total_cost == 0.005
|
||||
assert result3.total_cost == 0.008
|
||||
|
||||
# When summed for billing: 0.010 + 0.005 + 0.008 = 0.023
|
||||
# This verifies the cost accumulation fix (using += instead of max())
|
||||
total_from_events = result1.total_cost + result2.total_cost + result3.total_cost
|
||||
assert total_from_events == 0.023
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 3: Token Accumulation Consistency
|
||||
# ============================================================================
|
||||
@pytest.mark.unit
|
||||
def test_tokens_and_costs_extracted_independently() -> None:
|
||||
"""Tokens and costs should be extracted independently per event."""
|
||||
event = {
|
||||
"type": "content_block_delta",
|
||||
"usage": {
|
||||
"input_tokens": 100,
|
||||
"output_tokens": 50,
|
||||
"cache_read_input_tokens": 30,
|
||||
"total_cost": 0.007,
|
||||
"input_cost": 0.005,
|
||||
"output_cost": 0.002
|
||||
}
|
||||
}
|
||||
result = annotate_event(event, None)
|
||||
|
||||
# All values should be extracted
|
||||
assert result.input_tokens == 100
|
||||
assert result.output_tokens == 50
|
||||
assert result.cache_read_input_tokens == 30
|
||||
assert result.total_cost == 0.007
|
||||
assert result.input_cost == 0.005
|
||||
assert result.output_cost == 0.002
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 4: Cost in Message vs Root
|
||||
# ============================================================================
|
||||
@pytest.mark.unit
|
||||
def test_cost_extracted_from_message_usage() -> None:
|
||||
"""Costs in message.usage should be extracted correctly."""
|
||||
event = {
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"usage": {
|
||||
"input_tokens": 100,
|
||||
"total_cost": 0.010,
|
||||
"input_cost": 0.008,
|
||||
"output_cost": 0.002
|
||||
}
|
||||
}
|
||||
}
|
||||
result = annotate_event(event, None)
|
||||
|
||||
assert result.input_tokens == 100
|
||||
assert result.total_cost == 0.010
|
||||
assert result.input_cost == 0.008
|
||||
assert result.output_cost == 0.002
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 5: Cost at Event Root Level
|
||||
# ============================================================================
|
||||
@pytest.mark.unit
|
||||
def test_cost_extracted_from_event_root() -> None:
|
||||
"""Costs at event root should be extracted (OpenRouter style)."""
|
||||
event = {
|
||||
"type": "content_block_delta",
|
||||
"usage": {"output_tokens": 25},
|
||||
"total_cost": 0.005, # ← At root level
|
||||
"cost": 0.005
|
||||
}
|
||||
result = annotate_event(event, None)
|
||||
|
||||
assert result.output_tokens == 25
|
||||
assert result.total_cost == 0.005
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 6: Cost Details in Event Root
|
||||
# ============================================================================
|
||||
@pytest.mark.unit
|
||||
def test_cost_details_extracted_from_event_root() -> None:
|
||||
"""cost_details at event root should be extracted correctly."""
|
||||
event = {
|
||||
"type": "message_delta",
|
||||
"cost_details": {
|
||||
"total_cost": 0.015,
|
||||
"input_cost": 0.010,
|
||||
"output_cost": 0.005
|
||||
}
|
||||
}
|
||||
result = annotate_event(event, None)
|
||||
|
||||
assert result.total_cost == 0.015
|
||||
assert result.input_cost == 0.010
|
||||
assert result.output_cost == 0.005
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 7: No Duplicated Dict Lookups
|
||||
# ============================================================================
|
||||
@pytest.mark.unit
|
||||
def test_annotate_event_no_duplicate_lookups() -> None:
|
||||
"""Verify that dict lookups are not duplicated (fix for copy-paste error)."""
|
||||
# This is tested implicitly through proper extraction
|
||||
event = {
|
||||
"type": "message_delta",
|
||||
"cost_details": {
|
||||
"total_cost": 0.020,
|
||||
"input_cost": 0.015,
|
||||
"output_cost": 0.005
|
||||
}
|
||||
}
|
||||
result = annotate_event(event, None)
|
||||
|
||||
# Should extract each field exactly once, correctly
|
||||
assert result.total_cost == 0.020
|
||||
assert result.input_cost == 0.015
|
||||
assert result.output_cost == 0.005
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 8: Model Name Extraction
|
||||
# ============================================================================
|
||||
@pytest.mark.unit
|
||||
def test_model_extracted_from_event() -> None:
|
||||
"""Model name should be extracted from event."""
|
||||
event = {
|
||||
"type": "message_start",
|
||||
"message": {"model": "claude-3-5-sonnet"}
|
||||
}
|
||||
result = annotate_event(event, None)
|
||||
|
||||
assert result.model == "claude-3-5-sonnet"
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 9: Model Name Override
|
||||
# ============================================================================
|
||||
@pytest.mark.unit
|
||||
def test_model_name_override() -> None:
|
||||
"""Requested model should override actual model in event."""
|
||||
event = {
|
||||
"type": "message_start",
|
||||
"message": {"model": "actual-model"}
|
||||
}
|
||||
# Override with requested_model
|
||||
result = annotate_event(event, requested_model="alias-model")
|
||||
|
||||
# Event should be modified to use requested_model
|
||||
assert event["message"]["model"] == "alias-model" # type: ignore[index]
|
||||
assert result.model == "alias-model"
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 10: SSE Encoding
|
||||
# ============================================================================
|
||||
@pytest.mark.unit
|
||||
def test_sse_bytes_encoded() -> None:
|
||||
"""Event should be encoded as SSE bytes."""
|
||||
event = {
|
||||
"type": "content_block_start",
|
||||
"index": 0,
|
||||
"content_block": {"type": "text", "text": ""}
|
||||
}
|
||||
result = annotate_event(event, None)
|
||||
|
||||
assert result.sse_bytes is not None
|
||||
assert isinstance(result.sse_bytes, bytes)
|
||||
assert b"event: content_block_start" in result.sse_bytes or b"data:" in result.sse_bytes
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 11: Missing Cost Fields Default to Zero
|
||||
# ============================================================================
|
||||
@pytest.mark.unit
|
||||
def test_missing_cost_fields_default_to_zero() -> None:
|
||||
"""When cost fields are missing, should default to 0.0."""
|
||||
event = {
|
||||
"type": "content_block_delta",
|
||||
"usage": {"output_tokens": 25}
|
||||
# No cost fields
|
||||
}
|
||||
result = annotate_event(event, None)
|
||||
|
||||
assert result.total_cost == 0.0
|
||||
assert result.input_cost == 0.0
|
||||
assert result.output_cost == 0.0
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 12: Cache Tokens in Events
|
||||
# ============================================================================
|
||||
@pytest.mark.unit
|
||||
def test_cache_tokens_extracted_from_event() -> None:
|
||||
"""Cache tokens should be extracted from event usage."""
|
||||
event = {
|
||||
"type": "message_delta",
|
||||
"usage": {
|
||||
"input_tokens": 100,
|
||||
"output_tokens": 50,
|
||||
"cache_read_input_tokens": 200,
|
||||
"cache_creation_input_tokens": 0
|
||||
}
|
||||
}
|
||||
result = annotate_event(event, None)
|
||||
|
||||
assert result.input_tokens == 100
|
||||
assert result.output_tokens == 50
|
||||
assert result.cache_read_input_tokens == 200
|
||||
assert result.cache_creation_input_tokens == 0
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 13: Malformed Cost Values
|
||||
# ============================================================================
|
||||
@pytest.mark.unit
|
||||
def test_malformed_cost_values_coerced() -> None:
|
||||
"""Malformed cost values should be coerced to 0.0."""
|
||||
event = {
|
||||
"type": "message_delta",
|
||||
"usage": {
|
||||
"input_tokens": 100,
|
||||
"output_tokens": 50,
|
||||
"total_cost": "invalid", # ← Non-numeric
|
||||
}
|
||||
}
|
||||
result = annotate_event(event, None)
|
||||
|
||||
# Invalid cost should default to 0.0
|
||||
assert result.total_cost == 0.0
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 14: Negative Cost Values
|
||||
# ============================================================================
|
||||
@pytest.mark.unit
|
||||
def test_negative_cost_values_clamped() -> None:
|
||||
"""Negative cost values should be clamped to 0.0."""
|
||||
event = {
|
||||
"type": "message_delta",
|
||||
"usage": {
|
||||
"input_tokens": 100,
|
||||
"output_tokens": 50,
|
||||
"total_cost": -0.05 # ← Negative
|
||||
}
|
||||
}
|
||||
result = annotate_event(event, None)
|
||||
|
||||
# Negative cost should be clamped to 0.0
|
||||
assert result.total_cost == 0.0
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Test 15: Streaming Event Type Preserved
|
||||
# ============================================================================
|
||||
@pytest.mark.unit
|
||||
def test_event_type_preserved_in_sse() -> None:
|
||||
"""Event type should be preserved in SSE encoding."""
|
||||
event_types = ["message_start", "content_block_start", "content_block_delta", "message_delta"]
|
||||
for event_type in event_types:
|
||||
event = {"type": event_type}
|
||||
result = annotate_event(event, None)
|
||||
|
||||
assert result.event["type"] == event_type # type: ignore[index]
|
||||
# SSE should include the event type line if present
|
||||
if event_type:
|
||||
assert f"event: {event_type}".encode() in result.sse_bytes
|
||||
1313
tests/unit/test_messages_litellm_dispatch.py
Normal file
1313
tests/unit/test_messages_litellm_dispatch.py
Normal file
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user