Compare commits

...

174 Commits

Author SHA1 Message Date
9qeklajc
ee79a305ba revert 2026-05-30 17:20:16 +02:00
9qeklajc
ed2c8c9fe2 fix compose file 2026-05-30 17:00:57 +02:00
9qeklajc
9161256d31 Merge pull request #529 from Routstr/rip-08-lightning-invoice
better handling invoice payment
2026-05-30 16:39:11 +02:00
9qeklajc
1a8f407142 Merge pull request #531 from Routstr/remove-secp256-dep
remove secp256k1 dependency
2026-05-29 13:34:31 +02:00
9qeklajc
eddc070628 remove secp256k1 dependency 2026-05-28 21:09:36 +02:00
9qeklajc
f7bd250c97 better handling invoice payment 2026-05-28 20:09:48 +02:00
9qeklajc
29cfeaed9a Merge pull request #523 from Routstr/rip-08-lightning-invoice
add rop-08 lightning invoice support
2026-05-21 00:28:41 +02:00
9qeklajc
e2b46dab9f Merge pull request #522 from Routstr/fix-balance-info
use correct var to report balance info
2026-05-21 00:28:28 +02:00
9qeklajc
632244e54f add rop-08 lightning invoice support 2026-05-20 23:24:24 +02:00
9qeklajc
3251664513 use correct var to report balance info 2026-05-20 23:00:03 +02:00
9qeklajc
ee16cd0495 Merge pull request #518 from Routstr/tee-fixes
fix: tee GET passthrough and encoding fix
2026-05-19 22:31:19 +02:00
redshift
70ef3c357c fixed the encoding bug. 2026-05-19 10:53:59 +08:00
redshift
e4165f1dab passing through tee get requests 2026-05-19 10:53:59 +08:00
9qeklajc
41edae52b0 Merge pull request #513 from Routstr/add-provider-field-to-response
add provider field to response
2026-05-18 22:06:10 +02:00
9qeklajc
cb16da5543 Merge pull request #517 from Routstr/selinux-podman-fix
fix: add SELinux :z labels and user root to podman-compose volumes
2026-05-18 20:28:10 +02:00
redshift
d52b727bce fix: add SELinux :z labels and user root to podman-compose volumes
- Added :z (shared SELinux label) to all host bind-mount volumes
  so containers can write when SELinux is enforcing
- Set user: root on the ui service for compatibility with
  rootless podman's UID mapping
2026-05-19 01:44:51 +08:00
9qeklajc
6a2abd3439 Merge pull request #515 from Routstr/fix-value-not-set-issue
fix missing field initialization
2026-05-17 14:46:33 +02:00
9qeklajc
1e4ed75179 fix missing field initialization 2026-05-17 14:43:50 +02:00
9qeklajc
efb5719679 add provider field to response 2026-05-17 14:39:16 +02:00
9qeklajc
3a1902b2b5 Merge pull request #512 from Routstr/sidebar-display-logo
better logo display
2026-05-16 23:56:20 +02:00
9qeklajc
ba1978f830 better logo display 2026-05-16 23:52:03 +02:00
9qeklajc
2f51d62a46 Merge pull request #511 from Routstr/enforce-provider-fee-calculation
enforce fee calculation
2026-05-16 23:17:13 +02:00
9qeklajc
1ca344ba7b enforce fee calculation 2026-05-16 23:13:47 +02:00
9qeklajc
dda9b848ba Merge pull request #510 from Routstr/display-correct-version
display correct commit
2026-05-16 21:14:26 +02:00
9qeklajc
033d8c5a38 display correct commit 2026-05-16 19:10:45 +02:00
9qeklajc
e73cdffdc5 Merge pull request #508 from Routstr/display-official-version
display version
2026-05-16 15:46:43 +02:00
9qeklajc
e32b332eb5 Merge pull request #509 from Routstr/reduce-verbose-logs
less verbose logs and more precise
2026-05-16 15:41:20 +02:00
9qeklajc
89b84392a4 lint 2026-05-16 15:40:58 +02:00
9qeklajc
f37ff30598 display version 2026-05-16 15:35:16 +02:00
9qeklajc
547f45c185 less verbose logs and more precise 2026-05-16 15:29:41 +02:00
9qeklajc
a6e81a83cb Merge pull request #507 from Routstr/fix-provider-forwarding-upstream
Fix provider forwarding upstream
2026-05-15 20:45:19 +02:00
9qeklajc
e2143aa173 Revert "fix home nav"
This reverts commit 2eb257c2a6.
2026-05-14 15:51:20 +02:00
9qeklajc
2eb257c2a6 fix home nav 2026-05-14 15:47:59 +02:00
9qeklajc
490687bb71 Merge branch 'main' into fix-provider-forwarding-upstream 2026-05-14 15:40:52 +02:00
9qeklajc
966613847e update not found proxy 2026-05-14 15:40:49 +02:00
9qeklajc
d030d86f9a Merge pull request #506 from Routstr/fix-swapping
fix swapping
2026-05-14 14:32:39 +02:00
9qeklajc
d9ab46b3bb make sure when topup from untrusted mint to refund from primary mint 2026-05-14 13:52:25 +02:00
9qeklajc
6e596e3860 Merge pull request #504 from Routstr/do-not-create-empty-token
no key creation when refund
2026-05-13 22:48:16 +02:00
9qeklajc
3c13be20cb no key creation when refund 2026-05-13 22:34:24 +02:00
9qeklajc
9d5e14903b fix swapping 2026-05-13 21:49:03 +02:00
9qeklajc
9e33d3b100 Merge pull request #503 from Routstr/fix-docker-build
add missing path
2026-05-10 11:21:04 +02:00
9qeklajc
c88665f6d1 Merge pull request #500 from Routstr/enforce-loading-keysets
enforce load mint keyset
2026-05-10 11:20:53 +02:00
9qeklajc
a8157b3e2d add missing path 2026-05-10 11:18:27 +02:00
9qeklajc
a07e6723d9 Merge pull request #502 from Routstr/fix-docker-build
pin pnpm version
2026-05-10 10:44:44 +02:00
9qeklajc
5bc7b741bc pin pnpm version 2026-05-10 10:42:17 +02:00
9qeklajc
1db37cf084 Merge pull request #501 from Routstr/fix-docker-build
build fix
2026-05-10 10:20:13 +02:00
9qeklajc
6c5b103149 build fix 2026-05-10 10:18:07 +02:00
9qeklajc
27ec052c17 Merge pull request #499 from Routstr/emulate-claude-token-count
support /message/count_tokens endpoint
2026-05-10 08:45:58 +02:00
9qeklajc
00b813f6a8 clean up 2026-05-10 08:43:47 +02:00
9qeklajc
91e5198a94 Merge pull request #498 from Routstr/fix-docker-build
fix-docker-build
2026-05-10 08:43:02 +02:00
9qeklajc
27f1cc3c42 enforce load mint keyset 2026-05-10 08:32:19 +02:00
9qeklajc
a4d048f2b5 fix-docker-build 2026-05-10 08:31:05 +02:00
9qeklajc
752d4f3803 Merge pull request #496 from Routstr/fix-provider-forwarding-upstream
fix provder forwarded path & wrong html response
2026-05-09 14:47:35 +02:00
9qeklajc
164ed775c8 support /message/count_tokens endpoint 2026-05-09 14:43:03 +02:00
9qeklajc
9d905758ec added page not found redirection 2026-05-09 14:07:45 +02:00
9qeklajc
0d1cb66855 Merge pull request #494 from Routstr/payment-configuration
configure payment
2026-05-09 12:07:49 +02:00
9qeklajc
6e9932e0ac fix provder forwarded path & wrong html response 2026-05-09 12:05:51 +02:00
9qeklajc
59773e972b configure payment 2026-05-09 00:00:18 +02:00
9qeklajc
7fd0cdf987 Merge pull request #491 from Routstr/check-version-crrectly
check-version-correclty
2026-05-07 00:13:17 +02:00
9qeklajc
b1947d660a check-version-correclty 2026-05-07 00:08:54 +02:00
9qeklajc
d45a9edcba Merge pull request #490 from Routstr/add-cache-token-to-calculation
add-missing-cache-token-to-calculation
2026-05-06 22:54:34 +02:00
9qeklajc
4dbbb45240 add-missing-cache-token-to-calculation 2026-05-06 22:49:23 +02:00
9qeklajc
a286b5efa0 Merge pull request #488 from Routstr/match-on-forwarded-model-id-not-model-id
Match on forwarded model id not model
2026-05-06 21:54:17 +02:00
9qeklajc
26a0b04aef strict fallback 2026-05-06 00:46:42 +02:00
9qeklajc
f7ccc25a7f forward correct model id 2026-05-05 23:57:41 +02:00
9qeklajc
8dd0ece501 Merge pull request #487 from Routstr/use-full-image-for-ghci-
use full docker image for ghci
2026-05-05 22:55:42 +02:00
9qeklajc
10969e0719 update deployment 2026-05-05 22:40:47 +02:00
9qeklajc
b633455071 use full docker image for ghci 2026-05-05 22:15:13 +02:00
9qeklajc
e280d1f8b1 Merge pull request #479 from Routstr/litellm-integration-for-anthropic-messages-forwarding
Litellm integration for anthropic messages forwarding
2026-05-04 23:58:31 +02:00
9qeklajc
81ad9f604b clean up 2026-05-04 21:29:41 +02:00
9qeklajc
2345691176 fix gemini upstream claude func calls 2026-05-04 21:17:09 +02:00
9qeklajc
a7615bc827 improve gemini upstream to forward /messages endpoint correctly 2026-05-04 00:48:52 +02:00
9qeklajc
befdc5307e add logs for missing token usage 2026-05-03 22:40:06 +02:00
9qeklajc
145777ffd0 clean up 2026-05-03 15:35:49 +02:00
9qeklajc
37c2bea93d make sure to forward to the right upstream 2026-05-03 15:15:34 +02:00
9qeklajc
985e765285 fix gemini api 2026-05-02 21:24:27 +02:00
9qeklajc
a4b1330627 clean up impl. & simplify 2026-05-02 16:46:39 +02:00
9qeklajc
7b69284812 Merge branch 'main' into litellm-integration-for-anthropic-messages-forwarding 2026-05-01 23:09:52 +02:00
9qeklajc
7fb1da7b58 Merge pull request #483 from Routstr/more-logging-cleanup
remove and clean up redundant logs
2026-05-01 23:09:35 +02:00
9qeklajc
38f9923469 remove and clean up redundant logs 2026-05-01 23:07:43 +02:00
9qeklajc
5b986a3e15 Merge branch 'main' into litellm-integration-for-anthropic-messages-forwarding 2026-05-01 20:55:02 +02:00
9qeklajc
56d9ff6b3f Merge pull request #482 from Routstr/better-url-request-handling
handle urls correctly
2026-05-01 20:54:48 +02:00
9qeklajc
15b31e7c53 handle urls correctly 2026-05-01 17:08:05 +02:00
9qeklajc
d0dcc3219d Merge branch 'main' into litellm-integration-for-anthropic-messages-forwarding 2026-05-01 16:49:28 +02:00
9qeklajc
52db6fd308 Merge pull request #481 from Routstr/display-running-node-version
display commit if node not aligned with release tag
2026-05-01 16:49:15 +02:00
9qeklajc
ceca0e8efc display commit if node not aligned with release tag 2026-05-01 16:47:34 +02:00
9qeklajc
309bc873e6 Merge branch 'main' into litellm-integration-for-anthropic-messages-forwarding 2026-05-01 16:30:13 +02:00
9qeklajc
89109dc209 Merge pull request #480 from Routstr/clean-up-logging
do not logs redundant infos
2026-05-01 16:29:53 +02:00
9qeklajc
234ea19cad do not logs redundant infos 2026-05-01 16:23:30 +02:00
9qeklajc
489b6eba14 revert 2026-04-28 01:01:23 +02:00
9qeklajc
66421cd17f normalize response 2026-04-28 00:44:22 +02:00
9qeklajc
cd7de3958c Merge pull request #478 from Routstr/missing-model-pricing
make sure to always emit sats cost
2026-04-28 00:32:32 +02:00
9qeklajc
8c0ac499ef make sure to always emit sats cost 2026-04-28 00:01:27 +02:00
9qeklajc
1d22155e05 fix: default cashu MintInfo Optional fields to None for v2 parsing 2026-04-26 23:03:33 +02:00
9qeklajc
ab2fb2afc5 chore: bump cashu to 0.20 for pydantic v2 compat 2026-04-26 23:03:33 +02:00
9qeklajc
884bae9fc4 fix: migrate FastAPI-bound BaseModels to pydantic v2 to fix login 2026-04-26 22:57:31 +02:00
9qeklajc
92d581573d test: add unit tests for litellm messages dispatch 2026-04-26 22:45:22 +02:00
9qeklajc
85a3d3adc0 feat: route /v1/messages via litellm when upstream lacks native support 2026-04-26 22:45:22 +02:00
9qeklajc
2d7f03b2ed fix: switch openai NOT_GIVEN to omit for openai 2.x compat 2026-04-26 22:42:39 +02:00
9qeklajc
37bce70f76 docs: rewrite messages-to-chat-completions plan for litellm approach 2026-04-26 22:34:04 +02:00
9qeklajc
038bc14757 chore: add litellm dependency 2026-04-26 22:34:04 +02:00
9qeklajc
287cbac5f9 Merge branch 'admin-token' into litellm-integration-for-anthropic-messages-forwarding 2026-04-26 22:33:20 +02:00
9qeklajc
28d91227af Merge pull request #475 from Routstr/admin-token
Admin token
2026-04-26 22:32:21 +02:00
9qeklajc
846d894a13 chore: switch bare pydantic imports to v1 shim for v2 compat 2026-04-26 22:30:20 +02:00
9qeklajc
a2db3e2d57 fmt 2026-04-26 22:19:30 +02:00
9qeklajc
e43ceb2e43 Merge pull request #477 from Routstr/model-refresh
enforce-models-refresh-from-upstream
2026-04-26 22:03:59 +02:00
9qeklajc
689a07f562 enforce-models-refresh-from-upstream 2026-04-26 21:41:38 +02:00
9qeklajc
da487a850e Merge branch 'main' into admin-token 2026-04-26 00:09:00 +02:00
9qeklajc
fa6d3c76d0 Merge pull request #474 from Routstr/fix-field-label
use correct field label
2026-04-26 00:05:13 +02:00
9qeklajc
ede1804d4b use correct field label 2026-04-25 23:57:24 +02:00
9qeklajc
3392e8d4cb Merge pull request #473 from Routstr/bump-release-version
release v0.4.3
2026-04-25 11:56:48 +02:00
9qeklajc
b5174d9753 release v0.4.3 2026-04-25 11:54:54 +02:00
9qeklajc
1f2ff8a99c added admin token 2026-04-25 11:18:58 +02:00
9qeklajc
0c60644ba2 Merge pull request #472 from Routstr/fix-revision
fix migration
2026-04-24 23:08:06 +02:00
9qeklajc
aca8d43a61 fix migration 2026-04-24 23:05:09 +02:00
9qeklajc
7fe4c1963b Merge pull request #470 from Routstr/add-dev-cut
update migration rev id
2026-04-24 15:21:32 +02:00
9qeklajc
9a0919f149 update migration rev id 2026-04-24 15:19:53 +02:00
9qeklajc
d69ab913d4 Merge pull request #419 from Routstr/add-dev-cut
add-dev-cut
2026-04-23 23:22:23 +02:00
9qeklajc
42b8c332df Merge pull request #467 from Routstr/fix-concurent-refund-with-adding-token-history
Fix concurent refund with adding token history
2026-04-22 23:45:38 +02:00
9qeklajc
aa682bf8ec fix test 2026-04-22 23:43:11 +02:00
9qeklajc
c53e72e80a update migration 2026-04-22 23:31:49 +02:00
9qeklajc
afd81aeca2 reduce retries 2026-04-22 23:14:45 +02:00
9qeklajc
4e03145323 clean up 2026-04-22 23:13:29 +02:00
9qeklajc
169686681f Merge branch 'main' into add-dev-cut 2026-04-22 23:03:35 +02:00
9qeklajc
fa7d2804bb add test 2026-04-22 23:00:04 +02:00
9qeklajc
ccabc5d06d Merge branch 'main' into fix-response-error-forwarding 2026-04-22 22:17:59 +02:00
9qeklajc
cb68227c88 fix race cond. test 2026-04-22 22:17:41 +02:00
9qeklajc
c72fc7dd56 Merge pull request #466 from Routstr/add-detailed-response
fix forwarding upstream error responses
2026-04-22 22:16:41 +02:00
9qeklajc
8c6d1f89dc fix forwarding upstream error responses 2026-04-22 22:12:44 +02:00
9qeklajc
1ab7e54bd7 improve history collection 2026-04-22 22:09:13 +02:00
9qeklajc
42dceb0cd6 make table scrollable 2026-04-22 21:57:29 +02:00
9qeklajc
05115c3387 add api-key history and fix race condition while topup 2026-04-22 21:50:22 +02:00
9qeklajc
c0c8cafd00 fix forwarding upstream error responses 2026-04-20 16:45:05 +02:00
9qeklajc
16d6d66d17 Merge pull request #460 from Routstr/fix-test-interface-ui
Fix test interface UI
2026-04-19 15:29:00 +02:00
9qeklajc
3681fc8aab no advanced testing for now 2026-04-19 15:22:47 +02:00
9qeklajc
d95c09e00c fix test model connection 2026-04-19 15:21:23 +02:00
9qeklajc
3cdef5ec17 Merge pull request #459 from Routstr/add-detailed-logging
add logging reason for bad not succeed requests
2026-04-18 00:33:39 +02:00
9qeklajc
58e1620347 add logging reason for bad not succeed requests 2026-04-18 00:30:19 +02:00
9qeklajc
42b5a5e6c8 Merge pull request #458 from Routstr/add-cost-usage-to-messages-endpoint
Add cost usage to messages endpoint
2026-04-16 17:36:18 +02:00
9qeklajc
d84d249f2c Merge branch 'main' into add-dev-cut
# Conflicts:
#	routstr/auth.py
2026-04-15 23:07:37 +02:00
9qeklajc
8b6393f794 clean up 2026-04-15 21:56:30 +02:00
9qeklajc
81146710ba revert 2026-04-15 21:32:33 +02:00
9qeklajc
25f39897d7 mirror logic from chat completion 2026-04-15 20:45:13 +02:00
9qeklajc
45bdfaee58 best effort to be compatible with legacy code 2026-04-15 20:18:56 +02:00
9qeklajc
a1ae6e94e9 match model when versioned 2026-04-15 20:05:36 +02:00
9qeklajc
eebcc67c85 add cost usages 2026-04-15 17:19:13 +02:00
9qeklajc
3f9e7f7728 Merge pull request #457 from Routstr/fix-ppq-model-fetching
fix ppq model fetching
2026-04-14 18:49:26 +02:00
9qeklajc
a5b5549edd fix ppq model fetching 2026-04-14 18:46:38 +02:00
9qeklajc
22eec0162c Merge pull request #456 from Routstr/fix-serializing-broken-chunks
make streaming serialization more robust
2026-04-14 00:45:48 +02:00
9qeklajc
9f55da9bb8 make streaming serialization more robust 2026-04-14 00:44:10 +02:00
9qeklajc
16dce9ea81 Merge pull request #455 from Routstr/check-primary-mint
enforce no swap when token from primary mint
2026-04-14 00:32:06 +02:00
9qeklajc
963ee04619 enforce no swap when token from primary mint 2026-04-14 00:29:54 +02:00
9qeklajc
b874b1f01c Merge pull request #454 from Routstr/response-id-should-be-set
upstream drops id field which breaks opencode flow
2026-04-13 23:27:41 +02:00
9qeklajc
b6cca3d3a0 upstream drops id field which breaks opencode flow 2026-04-13 23:11:07 +02:00
9qeklajc
86ebc84f4c Merge pull request #453 from Routstr/remove-hardcoded-fee
remove hardcoded deduction (legacy code)
2026-04-13 22:47:23 +02:00
9qeklajc
8a74c0543f remove hardcoded deduction (legacy code) 2026-04-13 22:41:04 +02:00
9qeklajc
3d16a9b988 Merge pull request #449 from Routstr/apikey-transactions
add apikey refund tracking
2026-04-12 22:38:40 +02:00
9qeklajc
40f98b99aa set collected state 2026-04-12 22:36:52 +02:00
9qeklajc
7fcae5b08d Merge pull request #451 from Routstr/feature/update-deployment-docs
Updated deployment docs
2026-04-12 11:23:10 +02:00
redshift
f50cb31749 Updated deployment docs 2026-04-11 22:15:50 +01:00
9qeklajc
85a1ea4b4c add apikey refund tracking 2026-04-10 15:34:19 +02:00
9qeklajc
0f0f8c40bf Merge pull request #447 from Routstr/436-update-model-id
add support to using custom model ids
2026-04-10 14:18:24 +02:00
9qeklajc
a40d224ee7 Merge pull request #446 from Routstr/improve-negative-balance-handling
Improve negative balance handling
2026-04-10 14:14:22 +02:00
9qeklajc
a637acd8f4 use db fixture 2026-04-09 23:08:01 +02:00
9qeklajc
58be0c7976 clean u 2026-04-09 20:42:36 +02:00
9qeklajc
d934f3eead clean up 2026-04-09 20:41:59 +02:00
9qeklajc
74b58e5fa3 add some tests and improve payment handling 2026-04-09 20:25:43 +02:00
9qeklajc
c1497e0cfe Merge pull request #444 from Routstr/add-model-id-to-validation-log
add model id to validation log
2026-04-09 01:13:07 +02:00
9qeklajc
a1018776e9 make sure no negative balance can happen 2026-04-08 00:16:55 +02:00
9qeklajc
685368bb0a add support to using custom model ids 2026-04-06 20:09:12 +02:00
9qeklajc
453337cb2c update default payout 2026-04-05 00:36:59 +02:00
9qeklajc
236854bfe4 update lightining address 2026-03-25 10:25:46 +01:00
9qeklajc
a7886c528f Merge branch 'main' into add-dev-cut 2026-03-25 10:24:55 +01:00
9qeklajc
a7b815b29f add-dev-cut 2026-03-23 20:07:24 +01:00
93 changed files with 13925 additions and 3344 deletions

View File

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

View File

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

View File

@@ -4,7 +4,7 @@ FROM node:23-alpine AS ui-builder
WORKDIR /app/ui
# Install pnpm
RUN corepack enable pnpm && corepack prepare pnpm@latest --activate
RUN corepack enable pnpm && corepack prepare pnpm@10.15.0 --activate
# Copy UI source
COPY ui/package.json ui/pnpm-lock.yaml* ./
@@ -16,35 +16,37 @@ ENV NEXT_TELEMETRY_DISABLED=1
RUN pnpm run build
# Stage 2: Build the Routstr Node
FROM ghcr.io/astral-sh/uv:python3.11-alpine AS runner
FROM ghcr.io/astral-sh/uv:python3.11-bookworm-slim AS runner
# Install system dependencies
RUN apk add --no-cache \
pkgconf \
build-base \
automake \
autoconf \
libtool \
m4 \
perl \
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 uv sync --no-dev --no-install-project
WORKDIR /app
# Copy the rest of the application (required for uv sync to find the package)
COPY . .
# Install dependencies including the specific secp256k1 branch
RUN uv add git+https://github.com/saschanaz/secp256k1-py.git#branch=upgrade060
RUN uv sync --no-dev
# 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 ["/app/.venv/bin/fastapi", "run", "routstr", "--host", "0.0.0.0"]
CMD ["/.venv/bin/fastapi", "run", "routstr", "--host", "0.0.0.0"]

View File

@@ -51,14 +51,26 @@ curl https://api.routstr.com/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

View File

@@ -8,8 +8,9 @@ services:
args:
# 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:

View File

@@ -396,6 +396,8 @@ Authorization: Bearer sk-...
}
```
`balance` is the spendable balance used by request admission.
### Check Balance
Get current wallet balance.

View File

@@ -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-...`.

View File

@@ -78,9 +78,15 @@ Which mints to accept payments from:
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
@@ -126,6 +132,8 @@ Use environment variables for:
| `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) |

View File

@@ -2,20 +2,62 @@
Production deployment guide for Routstr Provider nodes.
## All-in-One Docker Image (Preferred)
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.
### Quick Start
```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: 9qeklajc/routstr:latest
container_name: routstr
restart: unless-stopped
ports:
- "8000:8000"
volumes:
- 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 it:
```bash
docker compose up -d
```
---
## Docker Compose (Recommended)
For production, use Docker Compose with persistent storage and optional Tor support.
### Unified Setup (All-in-one)
To build and run the node with the UI integrated in a single container using the multi-stage build:
```bash
docker build -f Dockerfile.full -t routstr-full .
docker run -d -p 8000:8000 --env-file .env routstr-full
```
### Advanced Setup (Separated UI & Node)
Use the included `compose.yml` for a more flexible setup that separates the UI build process from the node execution. This is useful for development or when you want to manage Tor as a separate service.
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
@@ -184,20 +226,16 @@ docker compose up -d
## Building from Source
### Unified Image (UI + Node)
The easiest way to build everything from source into a single production-ready image:
### Using Docker Compose
The easiest way to build everything from source:
```bash
docker build -f Dockerfile.full -t routstr-full .
docker compose build
```
### Individual Components
If you prefer building them separately or using Docker Compose:
If you prefer building the node only (requires manual UI build first):
```bash
# Build using compose
docker compose build
# Or build the node only (requires manual UI build first)
docker build -t routstr-node .
```

View File

@@ -35,6 +35,7 @@ ADMIN_PASSWORD=mysecretpassword
# Node Identity
NAME="My AI Node"
DESCRIPTION="Fast access to models"
NSEC=yournsec
# Lightning Payouts
RECEIVE_LN_ADDRESS=yourname@wallet.com
@@ -43,32 +44,10 @@ RECEIVE_LN_ADDRESS=yourname@wallet.com
## 2. Start the Node
You can run the pre-built image directly:
The recommended way to run Routstr is using Docker Compose, which handles the node, the UI, and optional services like Tor.
```bash
docker run -d \
--name routstr \
-p 8000:8000 \
--env-file .env \
-v routstr-data:/app/data \
ghcr.io/routstr/proxy:latest
```
*Note: The pre-built image does not contain the UI. For the all-in-one experience with the Admin Dashboard, use the Build from Source instructions below.*
### Build from Source (Recommended)
If you want to build the node and UI yourself from source, use the unified Dockerfile:
```bash
git clone https://github.com/routstr/routstr-core.git
cd routstr-core
# Edit your .env with ADMIN_PASSWORD and API keys
cp .env.example .env
nano .env
docker build -f Dockerfile.full -t routstr-local .
docker run -d -p 8000:8000 --env-file .env --name routstr routstr-local
docker compose up -d
```
Verify it's running:
@@ -77,6 +56,15 @@ 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
```
---
## 3. Configure via Dashboard

View 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")

View File

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

View File

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

View 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")

View File

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

View 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

View File

@@ -1,6 +1,6 @@
[project]
name = "routstr"
version = "0.4.1"
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" }

View File

@@ -195,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(
@@ -217,6 +218,10 @@ 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:
_add_candidate(alias, model_to_use, upstream)
@@ -268,18 +273,19 @@ def create_model_mappings(
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 base_id not in unique_models:
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[base_id] = unique_model
unique_models[unique_key] = unique_model
try:
aliases = resolve_model_alias(
@@ -305,6 +311,10 @@ 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)
for alias in aliases:
_add_candidate(alias, model_to_use, upstream_for_override)
seen_model_provider.add(dedupe_key)
@@ -314,7 +324,26 @@ def create_model_mappings(
provider_map: dict[str, list["BaseUpstreamProvider"]] = {}
def alias_priority(model: "Model", alias: str) -> int:
"""Rank how strong the mapping of alias->model is."""
"""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

View File

@@ -7,11 +7,12 @@ from datetime import datetime
from typing import Optional
from fastapi import HTTPException
from sqlalchemy import case
from sqlalchemy.exc import IntegrityError
from sqlmodel import col, select, update
from .core import get_logger
from .core.db import ApiKey, AsyncSession
from .core.db import ApiKey, AsyncSession, accumulate_routstr_fee
from .core.settings import settings
from .payment.cost_calculation import (
CostData,
@@ -22,6 +23,13 @@ from .payment.cost_calculation import (
from .wallet import credit_balance, deserialize_token_from_string
logger = get_logger(__name__)
payments_logger = get_logger("routstr.payments")
# Routstr platform fee constants
ROUTSTR_FEE_PERCENT: float = 2.1
ROUTSTR_LN_ADDRESS: str = "npub130mznv74rxs032peqym6g3wqavh472623mt3z5w73xq9r6qqdufs7ql29s@npub.cash"
ROUTSTR_FEE_PAYOUT_INTERVAL_SECONDS: int = 900
ROUTSTR_FEE_DEFAULT_PAYOUT: int = 200
# TODO: implement prepaid api key (not like it was before)
# PREPAID_API_KEY = os.environ.get("PREPAID_API_KEY", None)
@@ -309,13 +317,13 @@ async def validate_bearer_key(
extra={"key_hash": hashed_key[:8] + "..."},
)
logger.info(
logger.debug(
"AUTH: About to call credit_balance",
extra={"token_preview": bearer_key[:50]},
)
try:
msats = await credit_balance(bearer_key, new_key, session)
logger.info(
logger.debug(
"AUTH: credit_balance returned successfully", extra={"msats": msats}
)
except Exception as credit_error:
@@ -584,6 +592,18 @@ async def pay_for_request(
"total_requests": billing_key.total_requests,
},
)
payments_logger.info(
"RESERVE",
extra={
"event": "reserve",
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"cost_reserved": cost_per_request,
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
"total_spent": billing_key.total_spent,
},
)
return cost_per_request
@@ -635,6 +655,17 @@ async def revert_pay_for_request(
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
payments_logger.info(
"REVERT",
extra={
"event": "revert",
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"cost_reverted": cost_per_request,
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
},
)
return True
@@ -716,6 +747,17 @@ async def adjust_payment_for_tokens(
},
)
async def _accumulate_fee(total_cost_msats: int) -> None:
if total_cost_msats > 0 and ROUTSTR_FEE_PERCENT > 0:
fee_msats = math.ceil(total_cost_msats * ROUTSTR_FEE_PERCENT / 100)
try:
await accumulate_routstr_fee(session, fee_msats)
except Exception as e:
logger.warning(
"Failed to accumulate Routstr fee",
extra={"error": str(e), "fee_msats": fee_msats},
)
match await calculate_cost(response_data, deducted_max_cost, session):
case MaxCostData() as cost:
logger.debug(
@@ -728,11 +770,32 @@ async def adjust_payment_for_tokens(
},
)
# Finalize by releasing reservation and charging max cost
if billing_key.reserved_balance < deducted_max_cost:
logger.error(
"reserved_balance below deducted_max_cost before MaxCost finalization — clamping to 0",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"reserved_balance": billing_key.reserved_balance,
"deducted_max_cost": deducted_max_cost,
"total_cost_msats": cost.total_msats,
"balance": billing_key.balance,
"total_spent": billing_key.total_spent,
"model": model,
},
)
safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
finalize_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.values(
reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost,
reserved_balance=safe_reserved,
balance=col(ApiKey.balance) - cost.total_msats,
total_spent=col(ApiKey.total_spent) + cost.total_msats,
)
@@ -741,13 +804,17 @@ async def adjust_payment_for_tokens(
# Also update total_spent and reserved_balance on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(
total_spent=col(ApiKey.total_spent) + cost.total_msats,
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
reserved_balance=child_safe_reserved,
)
)
await session.exec(child_stmt) # type: ignore[call-overload]
@@ -782,6 +849,24 @@ async def adjust_payment_for_tokens(
"model": model,
},
)
await _accumulate_fee(cost.total_msats)
payments_logger.info(
"FINALIZE",
extra={
"event": "finalize",
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"model": model,
"cost_reserved": deducted_max_cost,
"cost_charged": cost.total_msats,
"input_tokens": cost.input_tokens,
"output_tokens": cost.output_tokens,
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
"total_spent": billing_key.total_spent,
"finalize_type": "max_cost",
},
)
return cost.dict()
case CostData() as cost:
@@ -815,12 +900,32 @@ async def adjust_payment_for_tokens(
"model": model,
},
)
if billing_key.reserved_balance < deducted_max_cost:
logger.error(
"reserved_balance below deducted_max_cost on exact-cost finalization — clamping to 0",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"reserved_balance": billing_key.reserved_balance,
"deducted_max_cost": deducted_max_cost,
"total_cost_msats": total_cost_msats,
"balance": billing_key.balance,
"total_spent": billing_key.total_spent,
"model": model,
},
)
exact_safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
finalize_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.values(
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
reserved_balance=exact_safe_reserved,
balance=col(ApiKey.balance) - total_cost_msats,
total_spent=col(ApiKey.total_spent) + total_cost_msats,
)
@@ -829,13 +934,17 @@ async def adjust_payment_for_tokens(
# Also update total_spent and reserved_balance on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_exact_safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(
total_spent=col(ApiKey.total_spent) + total_cost_msats,
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
reserved_balance=child_exact_safe_reserved,
)
)
await session.exec(child_stmt) # type: ignore[call-overload]
@@ -844,44 +953,56 @@ async def adjust_payment_for_tokens(
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
return cost.dict()
# this should never happen why do we handle this???
if cost_difference > 0:
# Need to charge more than reserved, finalize by releasing reservation and charging total
logger.info(
"Additional charge required for token usage",
await _accumulate_fee(total_cost_msats)
payments_logger.info(
"FINALIZE",
extra={
"event": "finalize",
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"additional_charge": cost_difference,
"current_balance": billing_key.balance,
"sufficient_balance": billing_key.balance >= cost_difference,
"model": model,
"cost_reserved": deducted_max_cost,
"cost_charged": total_cost_msats,
"input_tokens": cost.input_tokens,
"output_tokens": cost.output_tokens,
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
"total_spent": billing_key.total_spent,
"finalize_type": "exact",
},
)
return cost.dict()
# actual cost exceeded discounted reservation (due to tolerance_percentage)
if cost_difference > 0:
# Always release the reservation and charge min(actual_cost, balance).
# Using a CASE expression makes this a single atomic UPDATE — no
# multi-level fallback needed and balance can never go negative.
chargeable = case(
(col(ApiKey.balance) >= total_cost_msats, total_cost_msats),
else_=col(ApiKey.balance),
)
finalize_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.where(col(ApiKey.reserved_balance) >= deducted_max_cost)
.values(
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
balance=col(ApiKey.balance) - total_cost_msats,
total_spent=col(ApiKey.total_spent) + total_cost_msats,
reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost,
balance=col(ApiKey.balance) - chargeable,
total_spent=col(ApiKey.total_spent) + chargeable,
)
)
result = await session.exec(finalize_stmt) # type: ignore[call-overload]
# Also update total_spent and reserved_balance on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.where(col(ApiKey.reserved_balance) >= deducted_max_cost)
.values(
total_spent=col(ApiKey.total_spent) + total_cost_msats,
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost,
total_spent=col(ApiKey.total_spent) + min(billing_key.balance, total_cost_msats),
)
)
await session.exec(child_stmt) # type: ignore[call-overload]
@@ -889,11 +1010,10 @@ async def adjust_payment_for_tokens(
await session.commit()
if result.rowcount:
cost.total_msats = total_cost_msats
await session.refresh(billing_key)
if billing_key.hashed_key != key.hashed_key:
await session.refresh(key)
cost.total_msats = total_cost_msats
logger.info(
"Finalized payment with additional charge",
extra={
@@ -904,9 +1024,29 @@ async def adjust_payment_for_tokens(
"model": model,
},
)
await _accumulate_fee(total_cost_msats)
payments_logger.info(
"FINALIZE",
extra={
"event": "finalize",
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"model": model,
"cost_reserved": deducted_max_cost,
"cost_charged": total_cost_msats,
"input_tokens": cost.input_tokens,
"output_tokens": cost.output_tokens,
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
"total_spent": billing_key.total_spent,
"finalize_type": "overrun",
},
)
else:
# Guard fired: reservation was already released by a concurrent
# finalization for this key. Nothing left to do.
logger.warning(
"Failed to finalize additional charge - releasing reservation",
"Finalization skipped - reservation already released",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
@@ -914,7 +1054,6 @@ async def adjust_payment_for_tokens(
"model": model,
},
)
await release_reservation_only()
else:
# Refund some of the base cost
refund = abs(cost_difference)
@@ -929,12 +1068,33 @@ async def adjust_payment_for_tokens(
},
)
if billing_key.reserved_balance < deducted_max_cost:
logger.error(
"reserved_balance below deducted_max_cost on refund finalization — clamping to 0",
extra={
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"reserved_balance": billing_key.reserved_balance,
"deducted_max_cost": deducted_max_cost,
"total_cost_msats": total_cost_msats,
"refund_amount": refund,
"balance": billing_key.balance,
"total_spent": billing_key.total_spent,
"model": model,
},
)
refund_safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
refund_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == billing_key.hashed_key)
.values(
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
reserved_balance=refund_safe_reserved,
balance=col(ApiKey.balance) - total_cost_msats,
total_spent=col(ApiKey.total_spent) + total_cost_msats,
)
@@ -943,13 +1103,17 @@ async def adjust_payment_for_tokens(
# Also update total_spent and reserved_balance on the child key if it's different
if billing_key.hashed_key != key.hashed_key:
child_refund_safe_reserved = case(
(col(ApiKey.reserved_balance) >= deducted_max_cost,
col(ApiKey.reserved_balance) - deducted_max_cost),
else_=0,
)
child_stmt = (
update(ApiKey)
.where(col(ApiKey.hashed_key) == key.hashed_key)
.values(
total_spent=col(ApiKey.total_spent) + total_cost_msats,
reserved_balance=col(ApiKey.reserved_balance)
- deducted_max_cost,
reserved_balance=child_refund_safe_reserved,
)
)
await session.exec(child_stmt) # type: ignore[call-overload]
@@ -986,6 +1150,25 @@ async def adjust_payment_for_tokens(
"model": model,
},
)
await _accumulate_fee(total_cost_msats)
payments_logger.info(
"FINALIZE",
extra={
"event": "finalize",
"key_hash": key.hashed_key[:8] + "...",
"billing_key_hash": billing_key.hashed_key[:8] + "...",
"model": model,
"cost_reserved": deducted_max_cost,
"cost_charged": total_cost_msats,
"refunded": refund,
"input_tokens": cost.input_tokens,
"output_tokens": cost.output_tokens,
"balance": billing_key.balance,
"reserved_balance": billing_key.reserved_balance,
"total_spent": billing_key.total_spent,
"finalize_type": "refund",
},
)
return cost.dict()

View File

@@ -7,10 +7,16 @@ from typing import Annotated, NoReturn
from fastapi import APIRouter, Depends, Header, HTTPException
from fastapi.responses import JSONResponse
from pydantic import BaseModel
from sqlmodel import select
from sqlmodel import col, select, update
from .auth import get_billing_key, validate_bearer_key
from .core.db import ApiKey, AsyncSession, CashuTransaction, get_session
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
@@ -39,7 +45,7 @@ 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.balance,
"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,
@@ -205,6 +211,38 @@ async def _refund_cache_set(authorization: str, value: dict[str, str]) -> None:
_refund_cache[key] = (expiry, value)
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 | None, Header()] = None,
@@ -256,7 +294,12 @@ async def refund_wallet_endpoint(
)
bearer_value: str = authorization[7:]
key: ApiKey = await validate_bearer_key(bearer_value, session)
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 key.total_balance <= 0:
if cached := await _refund_cache_get(bearer_value):
@@ -286,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}
@@ -310,41 +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)
previous_reserved_balance = key.reserved_balance
key.balance = 0
key.reserved_balance = 0
session.add(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": previous_reserved_balance,
"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:
@@ -502,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)

View File

@@ -5,7 +5,8 @@ from datetime import datetime, timezone
from pathlib import Path
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from pydantic import BaseModel
from pydantic import BaseModel, RootModel
from pydantic.v1 import ValidationError as PydanticValidationError
from sqlmodel import select
from ..payment.models import _row_to_model, list_models
@@ -20,6 +21,8 @@ from ..wallet import (
from .db import (
ApiKey,
CashuTransaction,
CliToken,
LightningInvoice,
ModelRow,
UpstreamProviderRow,
create_session,
@@ -38,12 +41,27 @@ ADMIN_SESSION_DURATION = 3600
MAX_USAGE_ANALYTICS_HOURS = 365 * 24
def require_admin_api(request: Request) -> None:
async def require_admin_api(request: Request) -> None:
auth_header = request.headers.get("Authorization")
if auth_header and auth_header.startswith("Bearer "):
token = auth_header.split(" ", 1)[1]
expiry = admin_sessions.get(token)
if expiry and expiry > int(datetime.now(timezone.utc).timestamp()):
if not auth_header or not auth_header.startswith("Bearer "):
raise HTTPException(status_code=403, detail="Unauthorized")
token = auth_header.split(" ", 1)[1]
now_ts = int(datetime.now(timezone.utc).timestamp())
# 1) Short-lived session token (in-memory)
expiry = admin_sessions.get(token)
if expiry and expiry > now_ts:
return
# 2) Long-lived CLI token (DB-backed)
async with create_session() as session:
result = await session.exec(select(CliToken).where(CliToken.token == token))
cli_token = result.first()
if cli_token and (cli_token.expires_at is None or cli_token.expires_at > now_ts):
cli_token.last_used_at = now_ts
session.add(cli_token)
await session.commit()
return
raise HTTPException(status_code=403, detail="Unauthorized")
@@ -126,8 +144,8 @@ async def get_settings(request: Request) -> dict:
return data
class SettingsUpdate(BaseModel):
__root__: dict[str, object]
class SettingsUpdate(RootModel[dict[str, object]]):
pass
class PasswordUpdate(BaseModel):
@@ -138,14 +156,19 @@ class PasswordUpdate(BaseModel):
@admin_router.patch("/api/settings", dependencies=[Depends(require_admin_api)])
async def update_settings(request: Request, update: SettingsUpdate) -> dict:
# Remove sensitive fields from general settings update
settings_data = update.__root__.copy()
settings_data = update.root.copy()
sensitive_fields = ["admin_password", "upstream_api_key", "nsec"]
for field in sensitive_fields:
if field in settings_data:
del settings_data[field]
async with create_session() as session:
new_settings = await SettingsService.update(settings_data, session)
try:
async with create_session() as session:
new_settings = await SettingsService.update(settings_data, session)
except PydanticValidationError as e:
# Surface validation issues (e.g. non-positive payout amounts)
# as a clean 400 instead of a 500.
raise HTTPException(status_code=400, detail=e.errors()) from e
data = new_settings.dict()
if "upstream_api_key" in data:
data["upstream_api_key"] = "[REDACTED]" if data["upstream_api_key"] else ""
@@ -242,6 +265,73 @@ async def admin_logout(request: Request) -> dict[str, object]:
return {"ok": True}
# ─── CLI Tokens (long-lived bearer tokens for CLI/agent use) ───
class CliTokenCreate(BaseModel):
name: str
expires_in_days: int | None = None
@admin_router.get("/api/cli-tokens", dependencies=[Depends(require_admin_api)])
async def list_cli_tokens() -> list[dict[str, object]]:
async with create_session() as session:
result = await session.exec(select(CliToken))
tokens = result.all()
return [
{
"id": t.id,
"name": t.name,
"token_preview": f"{t.token[:8]}...{t.token[-4:]}",
"created_at": t.created_at,
"last_used_at": t.last_used_at,
"expires_at": t.expires_at,
}
for t in tokens
]
@admin_router.post("/api/cli-tokens", dependencies=[Depends(require_admin_api)])
async def create_cli_token(payload: CliTokenCreate) -> dict[str, object]:
name = (payload.name or "").strip()
if not name:
raise HTTPException(status_code=400, detail="Name is required")
raw_token = secrets.token_urlsafe(32)
expires_at: int | None = None
if payload.expires_in_days is not None and payload.expires_in_days > 0:
expires_at = int(datetime.now(timezone.utc).timestamp()) + (
payload.expires_in_days * 86400
)
async with create_session() as session:
cli_token = CliToken(token=raw_token, name=name, expires_at=expires_at)
session.add(cli_token)
await session.commit()
await session.refresh(cli_token)
return {
"id": cli_token.id,
"name": cli_token.name,
"token": raw_token, # full token returned only on creation
"created_at": cli_token.created_at,
"expires_at": cli_token.expires_at,
}
@admin_router.delete(
"/api/cli-tokens/{token_id}", dependencies=[Depends(require_admin_api)]
)
async def revoke_cli_token(token_id: str) -> dict[str, object]:
async with create_session() as session:
cli_token = await session.get(CliToken, token_id)
if not cli_token:
raise HTTPException(status_code=404, detail="Token not found")
await session.delete(cli_token)
await session.commit()
return {"ok": True, "deleted_id": token_id}
class WithdrawRequest(BaseModel):
amount: int
mint_url: str | None = None
@@ -295,6 +385,7 @@ class ModelCreate(BaseModel):
canonical_slug: str | None = None
alias_ids: list[str] | None = None
enabled: bool = True
forwarded_model_id: str | None = None
@admin_router.post(
@@ -339,6 +430,7 @@ async def upsert_provider_model(
json.dumps(payload.alias_ids) if payload.alias_ids else None
)
existing_row.enabled = payload.enabled
existing_row.forwarded_model_id = payload.forwarded_model_id or payload.id
session.add(existing_row)
await session.commit()
@@ -371,6 +463,7 @@ async def upsert_provider_model(
),
upstream_provider_id=provider_id,
enabled=payload.enabled,
forwarded_model_id=payload.forwarded_model_id or payload.id,
)
session.add(row)
await session.commit()
@@ -1332,41 +1425,102 @@ async def get_transactions_api(
type: str | None = None,
status: str | None = None,
search: str | None = None,
limit: int = 100,
source: str | None = None,
limit: int = 50,
offset: int = 0,
) -> dict:
async with create_session() as session:
from sqlmodel import col
from sqlmodel import col, func
stmt = select(CashuTransaction)
base = select(CashuTransaction)
if type:
stmt = stmt.where(CashuTransaction.type == type)
base = base.where(CashuTransaction.type == type)
if source:
if source == "x-cashu":
base = base.where(
(CashuTransaction.source == "x-cashu")
| (CashuTransaction.source == None) # noqa: E711
)
else:
base = base.where(CashuTransaction.source == source)
if status:
if status == "collected":
stmt = stmt.where(CashuTransaction.collected == True) # noqa: E712
base = base.where(CashuTransaction.collected == True) # noqa: E712
elif status == "swept":
stmt = stmt.where(CashuTransaction.swept == True) # noqa: E712
base = base.where(CashuTransaction.swept == True) # noqa: E712
elif status == "pending":
stmt = stmt.where(
base = base.where(
CashuTransaction.collected == False, # noqa: E712
CashuTransaction.swept == False, # noqa: E712
)
if search:
search_pattern = f"%{search}%"
stmt = stmt.where(
base = base.where(
(col(CashuTransaction.id).like(search_pattern))
| (col(CashuTransaction.token).like(search_pattern))
| (col(CashuTransaction.request_id).like(search_pattern))
| (col(CashuTransaction.api_key_hashed_key).like(search_pattern))
)
stmt = stmt.order_by(col(CashuTransaction.created_at).desc()).limit(limit)
count_result = await session.exec(
select(func.count()).select_from(base.subquery())
)
total = count_result.one()
stmt = base.order_by(col(CashuTransaction.created_at).desc()).offset(offset).limit(limit)
results = await session.exec(stmt)
transactions = results.all()
return {
"transactions": [tx.dict() for tx in transactions],
"total": len(transactions),
"total": total,
}
@admin_router.get(
"/api/lightning-invoices", dependencies=[Depends(require_admin_api)]
)
async def get_lightning_invoices_api(
status: str | None = None,
purpose: str | None = None,
search: str | None = None,
limit: int = 50,
offset: int = 0,
) -> dict:
async with create_session() as session:
from sqlmodel import col, func
base = select(LightningInvoice)
if status:
base = base.where(LightningInvoice.status == status)
if purpose:
base = base.where(LightningInvoice.purpose == purpose)
if search:
pattern = f"%{search}%"
base = base.where(
(col(LightningInvoice.id).like(pattern))
| (col(LightningInvoice.bolt11).like(pattern))
| (col(LightningInvoice.payment_hash).like(pattern))
| (col(LightningInvoice.api_key_hash).like(pattern))
)
count_result = await session.exec(
select(func.count()).select_from(base.subquery())
)
total = count_result.one()
stmt = (
base.order_by(col(LightningInvoice.created_at).desc())
.offset(offset)
.limit(limit)
)
results = await session.exec(stmt)
invoices = results.all()
return {
"invoices": [inv.dict() for inv in invoices],
"total": total,
}

View File

@@ -10,8 +10,9 @@ 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
@@ -78,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
@@ -105,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")
@@ -150,6 +154,16 @@ class CashuTransaction(SQLModel, table=True): # type: ignore
)
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(
@@ -161,6 +175,8 @@ async def store_cashu_transaction(
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:
@@ -173,6 +189,8 @@ async def store_cashu_transaction(
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()
@@ -212,6 +230,66 @@ class UpstreamProviderRow(SQLModel, table=True): # type: ignore
)
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:
@@ -327,6 +405,17 @@ def run_migrations() -> None:
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")

View File

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

View File

@@ -41,14 +41,23 @@ 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")
@@ -261,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,
@@ -270,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},
@@ -277,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,

View File

@@ -1,5 +1,4 @@
import asyncio
import os
from contextlib import asynccontextmanager
from pathlib import Path
from typing import AsyncGenerator
@@ -9,9 +8,12 @@ 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 ..lightning import lightning_router, periodic_invoice_watcher
from ..nostr import (
announce_provider,
providers_cache_refresher,
@@ -22,24 +24,22 @@ 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 ..upstream.auto_topup import periodic_auto_topup
from ..wallet import periodic_payout, periodic_refund_sweep
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.4.1-{os.getenv('VERSION_SUFFIX')}"
else:
__version__ = "0.4.1"
@asynccontextmanager
async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
@@ -56,8 +56,14 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, 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()
@@ -102,8 +108,11 @@ 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())
@@ -115,6 +124,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
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
@@ -152,6 +163,10 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, 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 = []
@@ -177,6 +192,10 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, 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)
@@ -188,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)
@@ -234,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",
)
@@ -242,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"
@@ -347,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"
@@ -369,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)

View File

@@ -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,56 +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
# 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),
"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),
},
)
@@ -82,36 +92,32 @@ 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),
},
)
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),
"error": str(e),
"error_type": type(e).__name__,

55
routstr/core/not_found.py Normal file
View 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)

View File

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

81
routstr/core/version.py Normal file
View 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"]

View File

@@ -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,15 +20,29 @@ 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)$")
purpose: str = Field(
default="create",
description="create or topup",
pattern="^(create|topup)$",
)
api_key: str | None = Field(
default=None, description="Required for topup operations"
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):
invoice_id: str
bolt11: str
@@ -64,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")
@@ -95,7 +113,7 @@ 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,
@@ -142,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:
@@ -274,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)

View File

@@ -26,7 +26,7 @@ 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:

View File

@@ -18,6 +18,10 @@ class CostData(BaseModel):
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):
@@ -29,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
@@ -51,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(
@@ -67,116 +78,58 @@ async def calculate_cost( # todo: can be sync
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"]
def parse_token_count(value: object) -> int:
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
# 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")
input_tokens = parse_token_count(usage_data.get("prompt_tokens", 0))
output_tokens = parse_token_count(usage_data.get("completion_tokens", 0))
input_tokens = (
input_tokens
if input_tokens != 0
else parse_token_count(usage_data.get("input_tokens", 0))
)
output_tokens = (
output_tokens
if output_tokens != 0
else parse_token_count(usage_data.get("output_tokens", 0))
)
input_tokens = (
input_tokens
if input_tokens != 0
else parse_token_count(response_data.get("usage", {}).get("input_tokens", 0))
)
output_tokens = (
output_tokens
if output_tokens != 0
else parse_token_count(response_data.get("usage", {}).get("output_tokens", 0))
)
usd_cost = 0.0
input_usd = 0.0
output_usd = 0.0
if "cost_details" in usage_data:
usd_cost = float(
usage_data["cost_details"].get("upstream_inference_cost", 0) or 0
)
input_usd = float(
usage_data["cost_details"].get("upstream_inference_prompt_cost", 0) or 0
)
output_usd = float(
usage_data["cost_details"].get("upstream_inference_completions_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
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
# 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)
input_msats = 0
output_msats = 0
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:
total_tokens = input_tokens + output_tokens
if total_tokens > 0:
input_ratio = input_tokens / total_tokens
input_msats = int(cost_in_msats * input_ratio)
output_msats = cost_in_msats - input_msats
else:
output_msats = cost_in_msats
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=0,
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,
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(
@@ -187,62 +140,35 @@ async def calculate_cost( # todo: can be sync
"model": response_data.get("model", "unknown"),
},
)
# Fall through to token-based calculation
if not settings.fixed_pricing:
response_model = response_data.get("model", "")
logger.debug(
"Using model-based pricing",
extra={"model": response_model},
)
# 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")
from ..proxy import get_model_instance
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
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(
@@ -252,12 +178,263 @@ async def calculate_cost( # todo: can be sync
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,
)
calc_input_msats = round(input_tokens / 1000 * MSATS_PER_1K_INPUT_TOKENS, 3)
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,
)
calc_output_msats = round(output_tokens / 1000 * MSATS_PER_1K_OUTPUT_TOKENS, 3)
token_based_cost = math.ceil(calc_input_msats + calc_output_msats)
# ============================================================================
# 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)
)
# 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)
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(
"Using cost from usage data/details",
extra={
"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=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(
@@ -265,8 +442,12 @@ async def calculate_cost( # todo: can be sync
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"),
@@ -281,4 +462,8 @@ async def calculate_cost( # todo: can be sync
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),
)

View File

@@ -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)
@@ -177,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:
@@ -329,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(
@@ -347,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
@@ -356,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()
@@ -402,6 +412,76 @@ 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("/v1/models/", include_in_schema=False)
@models_router.get("/models")
@@ -411,4 +491,10 @@ async def models(session: AsyncSession = Depends(get_session)) -> dict:
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}

View File

@@ -17,6 +17,7 @@ from .core.db import (
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,
@@ -69,7 +70,25 @@ 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) -> list[BaseUpstreamProvider] | None:
@@ -132,16 +151,72 @@ 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:
# 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)
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:
@@ -192,7 +267,15 @@ async def proxy(
)
except UpstreamError as e:
logger.warning(
f"Upstream {upstream.provider_type} failed (x-cashu): {e}"
"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
@@ -337,10 +420,14 @@ async def proxy(
)
logger.warning(
f"Upstream {upstream.provider_type} returned {response.status_code}, trying next provider",
"Upstream %s returned %s for model=%s, trying next provider",
upstream.provider_type,
response.status_code,
model_id,
extra={
"status_code": response.status_code,
"upstream": upstream.provider_type,
"provider": upstream.provider_type,
"model": model_id,
},
)
continue
@@ -348,16 +435,20 @@ async def proxy(
# 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",
"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,
"upstream_headers": response.headers
if hasattr(response, "headers")
else None,
},
)
return response
@@ -366,8 +457,16 @@ async def proxy(
except UpstreamError as e:
logger.warning(
f"Upstream {upstream.provider_type} failed: {e}",
extra={"retry": i < len(upstreams) - 1},
"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,
},
)
# If this was the last provider

View File

@@ -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__(

View File

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

File diff suppressed because it is too large Load Diff

View File

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

View 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",
)

View File

@@ -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__(

View File

@@ -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,248 +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
)
await session.refresh(key)
remaining_balance_msats = key.balance
openai_format_response["cost"] = cost_data
openai_format_response["cost"]["sats_cost"] = (
cost_data.get("total_msats", 0) // 1000
)
openai_format_response["cost"]["remaining_balance_msats"] = (
remaining_balance_msats
)
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()

View 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

View File

@@ -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__(

View File

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

View 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

View 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

View File

@@ -21,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,
@@ -184,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)},

View File

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

View File

@@ -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__(

View File

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

View File

@@ -18,6 +18,10 @@ class RoutstrUpstreamProvider(BaseUpstreamProvider):
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,
@@ -43,6 +47,13 @@ class RoutstrUpstreamProvider(BaseUpstreamProvider):
)
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"

View File

@@ -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__(

View File

@@ -1,16 +1,32 @@
import asyncio
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 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,6 +48,8 @@ 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)
@@ -40,17 +58,56 @@ async def recieve_token(
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
@@ -65,10 +122,11 @@ async def _calculate_swap_amount(
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
potential swap fees (melt fees) on the foreign mint.
melt fees and NUT-02 input fees on the foreign mint.
"""
if settings.primary_mint_unit == "sat":
receive_amount = amount_msat // 1000
@@ -95,10 +153,11 @@ async def _calculate_swap_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 * 1000
fee_msat = (fee_reserve + input_fees) * 1000
else:
fee_msat = fee_reserve
fee_msat = fee_reserve + input_fees
amount_msat_after_fee = amount_msat - fee_msat
@@ -108,13 +167,14 @@ async def _calculate_swap_amount(
minted_amount = int(amount_msat_after_fee)
if minted_amount <= 0:
raise ValueError(f"Fees ({fee_reserve} {token_unit}) exceed token amount")
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,
},
@@ -155,12 +215,27 @@ async def swap_to_primary_mint(
raise ValueError("Invalid unit")
primary_wallet = await get_wallet(settings.primary_mint, settings.primary_mint_unit)
# 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)
@@ -170,13 +245,15 @@ async def swap_to_primary_mint(
)
melt_quote = await token_wallet.melt_quote(mint_quote.request)
total_needed = melt_quote.amount + melt_quote.fee_reserve
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,
},
@@ -189,6 +266,7 @@ async def swap_to_primary_mint(
"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,
},
@@ -196,7 +274,7 @@ async def swap_to_primary_mint(
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})"
f"(amount: {melt_quote.amount} + fee: {melt_quote.fee_reserve} + input_fees: {input_fees})"
)
try:
@@ -227,19 +305,66 @@ async def swap_to_primary_mint(
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:
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
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",
@@ -265,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},
@@ -296,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},
)
@@ -398,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:
@@ -446,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,
@@ -457,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:
@@ -479,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,
@@ -535,7 +682,7 @@ async def periodic_refund_sweep() -> None:
except Exception as e:
error_msg = str(e).lower()
if "already spent" in error_msg:
refund.swept = True
refund.collected = True
session.add(refund)
logger.info(
"Refund already spent (client collected), marking swept",
@@ -559,6 +706,46 @@ async def periodic_refund_sweep() -> None:
)
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]

View File

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

View 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"
)

View 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)

View 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

View 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

View 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

View File

@@ -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
@@ -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(
@@ -394,6 +458,42 @@ async def test_refund_during_active_usage(
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
@pytest.mark.asyncio
async def test_mint_unavailability_handling(

View File

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

View File

@@ -1,11 +1,12 @@
import json
from unittest.mock import AsyncMock, MagicMock
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 CashuTransaction
from routstr.core.db import ApiKey, CashuTransaction
from routstr.wallet import credit_balance
def _make_cashu_tx(
@@ -29,6 +30,12 @@ def _exec_result(tx: CashuTransaction | None) -> MagicMock:
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"
@@ -114,3 +121,283 @@ async def test_refund_x_cashu_swept_raises_410() -> None:
)
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()

View 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

View 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

View 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/"
)

View 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

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,117 @@
"""Regression tests for the periodic upstream models refresh loop."""
from __future__ import annotations
import asyncio
import os
from typing import cast
from unittest.mock import AsyncMock
import pytest
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
os.environ.setdefault("UPSTREAM_API_KEY", "test")
from routstr.upstream.base import BaseUpstreamProvider # noqa: E402
class _FakeUpstream:
"""Minimal stand-in for BaseUpstreamProvider used by the refresh loop.
Only ``base_url`` (for error logging) and ``refresh_models_cache`` (the call
under test) are exercised; everything else stays unused.
"""
def __init__(self, name: str) -> None:
self.base_url = f"http://{name}"
self.refresh_models_cache = AsyncMock()
def _make_fake_upstream(name: str) -> BaseUpstreamProvider:
# The loop only uses duck-typed attributes — cast keeps the test type-clean
# without dragging in BaseUpstreamProvider's full constructor.
return cast(BaseUpstreamProvider, _FakeUpstream(name))
@pytest.mark.asyncio
async def test_refresh_loop_picks_up_providers_added_after_startup(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""If a provider is added after the loop starts (e.g. via reinitialize_upstreams),
the next loop iteration must refresh it. Previously the loop captured the upstream
list at startup and missed any later additions."""
from routstr.core.settings import settings as global_settings
from routstr.upstream.helpers import refresh_upstreams_models_periodically
# Tight interval so the test finishes quickly.
monkeypatch.setattr(
global_settings, "models_refresh_interval_seconds", 1, raising=False
)
initial_upstream = _make_fake_upstream("initial")
live_list: list[BaseUpstreamProvider] = [initial_upstream]
# Stub out the post-iteration sats-pricing refresh so the loop body has no DB deps.
async def _noop_pricing_refresh() -> None: # pragma: no cover - trivial stub
return None
monkeypatch.setattr(
"routstr.payment.models._update_sats_pricing_once",
_noop_pricing_refresh,
)
task = asyncio.create_task(
refresh_upstreams_models_periodically(lambda: live_list)
)
try:
# Wait for the first iteration to refresh the initial upstream.
for _ in range(40):
if initial_upstream.refresh_models_cache.await_count >= 1: # type: ignore[attr-defined]
break
await asyncio.sleep(0.05)
assert initial_upstream.refresh_models_cache.await_count >= 1, ( # type: ignore[attr-defined]
"loop did not refresh the initial upstream within the timeout"
)
# Simulate reinitialize_upstreams: replace the live list contents with new
# provider instances. The loop must observe the swap on its next tick.
new_upstream = _make_fake_upstream("added-after-startup")
live_list[:] = [new_upstream]
for _ in range(60):
if new_upstream.refresh_models_cache.await_count >= 1: # type: ignore[attr-defined]
break
await asyncio.sleep(0.05)
assert new_upstream.refresh_models_cache.await_count >= 1, ( # type: ignore[attr-defined]
"loop did not refresh the upstream added after startup — "
"regression: list snapshot captured at startup"
)
finally:
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
@pytest.mark.asyncio
async def test_refresh_loop_disabled_when_interval_non_positive(
monkeypatch: pytest.MonkeyPatch,
) -> None:
from routstr.core.settings import settings as global_settings
from routstr.upstream.helpers import refresh_upstreams_models_periodically
monkeypatch.setattr(
global_settings, "models_refresh_interval_seconds", 0, raising=False
)
upstream = _make_fake_upstream("never-refreshed")
# Loop must return immediately without ever touching the upstream.
await asyncio.wait_for(
refresh_upstreams_models_periodically(lambda: [upstream]),
timeout=1.0,
)
upstream.refresh_models_cache.assert_not_awaited() # type: ignore[attr-defined]

View File

@@ -0,0 +1,105 @@
from routstr.upstream.anthropic import AnthropicUpstreamProvider
from routstr.upstream.base import BaseUpstreamProvider
from routstr.upstream.openrouter import OpenRouterUpstreamProvider
def _make_provider(cls: type, provider_type: str) -> BaseUpstreamProvider:
p = cls(api_key="test_key")
assert p.provider_type == provider_type
return p
def test_apply_provider_field_direct_upstream() -> None:
"""For a direct upstream (no upstream-reported provider), the field
is just the provider_type string."""
p = _make_provider(AnthropicUpstreamProvider, "anthropic")
data: dict = {"id": "msg_1", "model": "claude-3-5-sonnet"}
p._apply_provider_field(data)
assert data["provider"] == "anthropic"
def test_apply_provider_field_openrouter_passthrough() -> None:
"""OpenRouter responses include an upstream ``provider`` string —
routstr should prefix with its own provider_type."""
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data: dict = {
"id": "gen-abc",
"model": "anthropic/claude-3.5-sonnet",
"provider": "Anthropic",
}
p._apply_provider_field(data)
assert data["provider"] == "openrouter:Anthropic"
def test_apply_provider_field_openrouter_no_upstream_provider() -> None:
"""If OpenRouter omits the provider field, fall back to provider_type."""
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data: dict = {"id": "gen-abc"}
p._apply_provider_field(data)
assert data["provider"] == "openrouter"
def test_apply_provider_field_strips_whitespace() -> None:
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data: dict = {"provider": " Fireworks "}
p._apply_provider_field(data)
assert data["provider"] == "openrouter:Fireworks"
def test_apply_provider_field_blank_upstream_treated_as_missing() -> None:
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data: dict = {"provider": " "}
p._apply_provider_field(data)
assert data["provider"] == "openrouter"
def test_apply_provider_field_non_string_upstream_treated_as_missing() -> None:
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
data: dict = {"provider": 42}
p._apply_provider_field(data)
assert data["provider"] == "openrouter"
def test_apply_provider_field_idempotent_for_direct_upstream() -> None:
"""Calling twice on a direct upstream payload should keep the same
value, not nest the prefix repeatedly."""
p = _make_provider(AnthropicUpstreamProvider, "anthropic")
data: dict = {}
p._apply_provider_field(data)
p._apply_provider_field(data)
assert data["provider"] == "anthropic:anthropic"
# Document current (deliberate) behavior: second pass treats the
# first-pass value as an upstream-reported provider. Callers should
# only invoke this once per chunk — guarded via the
# ``"provider" not in data`` checks in streaming paths.
def test_apply_provider_field_ignores_non_dict() -> None:
"""Lists / primitives must be skipped silently."""
p = _make_provider(AnthropicUpstreamProvider, "anthropic")
# Should not raise.
p._apply_provider_field([1, 2, 3]) # type: ignore[arg-type]
p._apply_provider_field("hello") # type: ignore[arg-type]
p._apply_provider_field(None) # type: ignore[arg-type]
def test_inject_cost_metadata_sets_provider() -> None:
"""``inject_cost_metadata`` is the unified injection point and must
also stamp the provider field."""
from unittest.mock import MagicMock
from routstr.core.db import ApiKey
p = _make_provider(OpenRouterUpstreamProvider, "openrouter")
key = MagicMock(spec=ApiKey)
key.balance = 1000
response_json: dict = {
"model": "anthropic/claude-3.5-sonnet",
"provider": "Anthropic",
"usage": {"prompt_tokens": 10, "completion_tokens": 5},
}
cost_data = {"total_msats": 2500, "total_usd": 0.0025}
p.inject_cost_metadata(response_json, cost_data, key)
assert response_json["provider"] == "openrouter:Anthropic"

View File

@@ -0,0 +1,70 @@
"""Tests for the app-level 404 handler in routstr.core.main."""
from __future__ import annotations
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from routstr.core import main as core_main
def _make_app() -> FastAPI:
app = FastAPI()
app.add_api_route(
"/{path:path}",
core_main.not_found_catch_all,
methods=["GET", "POST"],
include_in_schema=False,
)
return app
@pytest.mark.skipif(
core_main._NOT_FOUND_HTML is None,
reason="UI bundle (ui_out/404.html) not present in this environment",
)
def test_unknown_path_returns_html_404_for_browser() -> None:
client = TestClient(_make_app())
response = client.get("/some/random/page", headers={"accept": "text/html"})
assert response.status_code == 404
assert response.headers["content-type"].startswith("text/html")
assert "404" in response.text
def test_unknown_path_returns_json_404_for_api_client() -> None:
client = TestClient(_make_app())
response = client.get(
"/some/random/page", headers={"accept": "application/json"}
)
assert response.status_code == 404
assert response.headers["content-type"].startswith("application/json")
payload = response.json()
assert payload["error"]["type"] == "not_found"
assert payload["error"]["code"] == 404
assert "/some/random/page" in payload["error"]["message"]
def test_root_path_returns_404() -> None:
client = TestClient(_make_app())
response = client.get("/", headers={"accept": "application/json"})
assert response.status_code == 404
def test_json_returned_when_ui_html_missing(monkeypatch: pytest.MonkeyPatch) -> None:
from routstr.core import not_found as nf
monkeypatch.setattr(nf, "_NOT_FOUND_HTML", None)
client = TestClient(_make_app())
response = client.get("/some/random/page", headers={"accept": "text/html"})
assert response.status_code == 404
assert response.headers["content-type"].startswith("application/json")
def test_post_unknown_path_returns_json_even_for_browser() -> None:
client = TestClient(_make_app())
response = client.post("/some/random/page", headers={"accept": "text/html"})
assert response.status_code == 404
assert response.headers["content-type"].startswith("application/json")
payload = response.json()
assert payload["error"]["type"] == "not_found"

View File

@@ -1,11 +1,12 @@
import os
import pytest
from pydantic.v1 import ValidationError
from sqlalchemy.ext.asyncio import create_async_engine
from sqlmodel import text
from sqlmodel.ext.asyncio.session import AsyncSession
from routstr.core.settings import SettingsService
from routstr.core.settings import Settings, SettingsService
@pytest.mark.asyncio
@@ -46,6 +47,54 @@ async def test_settings_db_precedence_over_env() -> None:
assert again.enable_analytics_sharing is False
def test_payout_settings_have_sensible_defaults() -> None:
s = Settings()
assert s.min_payout_sat == 210
assert s.payout_interval_seconds == 900
@pytest.mark.parametrize(
"field,bad_value",
[
("min_payout_sat", 0),
("min_payout_sat", -1),
("payout_interval_seconds", 0),
("payout_interval_seconds", -10),
],
)
def test_payout_settings_reject_invalid_values(field: str, bad_value: int) -> None:
kwargs: dict[str, object] = {field: bad_value}
with pytest.raises(ValidationError):
Settings(**kwargs) # type: ignore[arg-type]
def test_payout_settings_accept_custom_positive_values() -> None:
s = Settings(min_payout_sat=500, payout_interval_seconds=60)
assert s.min_payout_sat == 500
assert s.payout_interval_seconds == 60
@pytest.mark.asyncio
async def test_payout_settings_persist_via_settings_service() -> None:
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
async with AsyncSession(engine, expire_on_commit=False) as session:
await SettingsService.initialize(session)
updated = await SettingsService.update(
{"min_payout_sat": 1000, "payout_interval_seconds": 300}, session
)
assert updated.min_payout_sat == 1000
assert updated.payout_interval_seconds == 300
@pytest.mark.asyncio
async def test_payout_settings_update_rejects_invalid() -> None:
engine = create_async_engine("sqlite+aiosqlite:///:memory:")
async with AsyncSession(engine, expire_on_commit=False) as session:
await SettingsService.initialize(session)
with pytest.raises(ValidationError):
await SettingsService.update({"min_payout_sat": 0}, session)
@pytest.mark.asyncio
async def test_settings_initialize_discards_unknown_keys() -> None:
engine = create_async_engine("sqlite+aiosqlite:///:memory:")

View File

@@ -0,0 +1,117 @@
import json
from collections.abc import AsyncGenerator
from unittest.mock import AsyncMock, MagicMock
import pytest
from routstr.core.db import ApiKey
from routstr.upstream.base import BaseUpstreamProvider
@pytest.mark.asyncio
async def test_stream_with_id_injection() -> None:
"""Test that stream_with_cost correctly injects IDs into complete JSON chunks but skips partials."""
provider = BaseUpstreamProvider(
base_url="https://api.example.com", api_key="test_key"
)
# Mock response with mixed chunks:
# 1. Complete JSON without ID
# 2. Partial JSON (should be passed through)
# 3. Complete JSON with ID (should be preserved or updated if requested_model is set)
# 4. [DONE] message
chunks = [
b'data: {"choices": [{"delta": {"content": "Hello"}}]}\n\n',
b'data: {"choices": [{"delta": {"content": "', # Partial
b'world"}}]}\n\n',
b'data: {"id": "existing-id", "choices": [{"delta": {"content": "!"}}]}\n\n',
b"data: [DONE]\n\n",
]
async def aiter_bytes() -> AsyncGenerator[bytes, None]:
for chunk in chunks:
yield chunk
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "text/event-stream"}
mock_response.aiter_bytes = aiter_bytes
key = MagicMock(spec=ApiKey)
key.hashed_key = "test_hash"
key.balance = 1000
background_tasks = MagicMock()
# We need to mock adjust_payment_for_tokens since it's called at the end
with MagicMock():
from routstr.upstream import base
# Mocking the module-level function used in the generator
base.adjust_payment_for_tokens = AsyncMock(
return_value={"total_usd": 0.1, "total_msats": 100}
)
# create_session() is used as an async context manager whose entered
# value exposes an awaitable .get(). Build a mock that behaves that
# way so the post-stream cost-chunk emission can run.
mock_session = MagicMock()
mock_session.get = AsyncMock(return_value=key)
mock_ctx = MagicMock()
mock_ctx.__aenter__ = AsyncMock(return_value=mock_session)
mock_ctx.__aexit__ = AsyncMock(return_value=None)
base.create_session = MagicMock(return_value=mock_ctx)
streaming_response = await provider.handle_streaming_chat_completion(
response=mock_response,
key=key,
max_cost_for_model=100,
background_tasks=background_tasks,
requested_model="test-model",
)
results = []
async for chunk in streaming_response.body_iterator:
results.append(chunk)
# Parse results
parsed_results = []
for r in results:
if isinstance(r, bytes) and r.startswith(b"data: "):
data = r[6:].decode().strip()
if data == "[DONE]":
parsed_results.append(data)
else:
try:
parsed_results.append(json.loads(data))
except (json.JSONDecodeError, UnicodeDecodeError):
parsed_results.append(
data
) # Keep as string if it failed to parse
# Verifications
# 1. First chunk should have an injected ID and the requested model
assert isinstance(parsed_results[0], dict)
assert "id" in parsed_results[0]
assert parsed_results[0]["id"].startswith("chatcmpl-")
assert parsed_results[0]["model"] == "test-model"
# 2. Second chunk was partial, should be passed as-is
# In current implementation, re.split(b"data: ", b'data: {...') gives ['', '{...']
# The first empty part is skipped. The second part is processed.
# Check that we have results
assert len(parsed_results) >= 4
# Find the chunk that was "existing-id"
id_chunk = next(
r
for r in parsed_results
if isinstance(r, dict)
and "choices" in r
and r["choices"][0]["delta"].get("content") == "!"
)
assert id_chunk["id"] == parsed_results[0]["id"]
assert id_chunk["model"] == "test-model"
# 4. [DONE] should be there
assert "[DONE]" in parsed_results

View File

@@ -0,0 +1,150 @@
"""Tests for ``BaseUpstreamProvider.forward_upstream_error_response``.
Upstream services (e.g. an Express server that doesn't expose ``/messages``)
sometimes return a non-JSON error body. The proxy must surface those errors
in a consistent JSON envelope so clients don't have to parse HTML.
"""
from __future__ import annotations
import json
from typing import Any
from unittest.mock import Mock
import httpx
import pytest
from routstr.upstream.base import BaseUpstreamProvider, _is_json_content_type
def _make_request(request_id: str = "req-123") -> Mock:
request = Mock(spec=["method", "state"])
request.method = "POST"
request.state = Mock()
request.state.request_id = request_id
return request
def _make_upstream_response(
*,
body: bytes,
status_code: int = 404,
content_type: str | None = "text/html",
extra_headers: dict[str, str] | None = None,
) -> httpx.Response:
headers: dict[str, str] = {}
if content_type is not None:
headers["content-type"] = content_type
if extra_headers:
headers.update(extra_headers)
return httpx.Response(status_code=status_code, headers=headers, content=body)
@pytest.fixture
def provider() -> BaseUpstreamProvider:
return BaseUpstreamProvider(
base_url="https://privateprovider.xyz", api_key="k", provider_fee=1.0
)
@pytest.mark.parametrize(
"content_type,expected",
[
("application/json", True),
("application/json; charset=utf-8", True),
("text/json", True),
("application/problem+json", True),
("application/vnd.api+json", True),
("text/html", False),
("text/html; charset=utf-8", False),
("text/plain", False),
("", False),
(None, False),
],
)
def test_is_json_content_type(content_type: str | None, expected: bool) -> None:
assert _is_json_content_type(content_type) is expected
@pytest.mark.asyncio
async def test_html_error_is_normalized_to_json_envelope(
provider: BaseUpstreamProvider,
) -> None:
html_body = (
b"<!DOCTYPE html><html><head><title>Error</title></head>"
b"<body><pre>Cannot POST /messages</pre></body></html>"
)
upstream = _make_upstream_response(body=html_body, status_code=404)
response = await provider.forward_upstream_error_response(
_make_request(), "v1/messages", upstream
)
assert response.status_code == 404
assert response.media_type == "application/json"
payload: dict[str, Any] = json.loads(bytes(response.body))
assert payload["error"]["type"] == "upstream_error"
assert payload["error"]["upstream_status"] == 404
assert payload["error"]["upstream_content_type"] == "text/html"
assert "Cannot POST /messages" in payload["error"]["upstream_body_preview"]
assert payload["request_id"] == "req-123"
# The upstream's text/html content-type must not survive — Response()
# sets the JSON content-type for us via media_type.
assert response.headers["content-type"].startswith("application/json")
@pytest.mark.asyncio
async def test_plain_text_error_is_normalized(
provider: BaseUpstreamProvider,
) -> None:
upstream = _make_upstream_response(
body=b"Service Unavailable", status_code=503, content_type="text/plain"
)
response = await provider.forward_upstream_error_response(
_make_request(), "v1/messages", upstream
)
assert response.status_code == 503
assert response.media_type == "application/json"
payload = json.loads(bytes(response.body))
assert payload["error"]["message"] == "Service Unavailable"
@pytest.mark.asyncio
async def test_empty_body_with_non_json_content_type_normalizes(
provider: BaseUpstreamProvider,
) -> None:
upstream = _make_upstream_response(
body=b"", status_code=502, content_type="text/html"
)
response = await provider.forward_upstream_error_response(
_make_request(), "v1/messages", upstream
)
assert response.status_code == 502
assert response.media_type == "application/json"
payload = json.loads(bytes(response.body))
assert payload["error"]["type"] == "upstream_error"
assert payload["error"]["upstream_body_preview"] is None
@pytest.mark.asyncio
async def test_json_error_body_is_passed_through_unchanged(
provider: BaseUpstreamProvider,
) -> None:
json_body = json.dumps(
{"error": {"message": "Invalid model", "type": "invalid_request_error"}}
).encode()
upstream = _make_upstream_response(
body=json_body, status_code=400, content_type="application/json"
)
response = await provider.forward_upstream_error_response(
_make_request(), "v1/messages", upstream
)
assert response.status_code == 400
assert bytes(response.body) == json_body
assert response.media_type == "application/json"

View File

@@ -0,0 +1,313 @@
"""Unit tests for the Gemini /v1/messages dispatch path.
The Gemini upstream needs special handling because its OpenAI-compat
surface rejects inbound ``functionCall`` parts that lack a
``thought_signature``. We bypass litellm + openai SDK at the wire layer
(see ``routstr/upstream/gemini_messages.py``) so we can inject Google's
documented dummy signature (``"skip_thought_signature_validator"``).
These tests cover the two pure helpers that drive the dispatcher:
* ``inject_thought_signatures`` — request-side injection
* ``_openai_chunks_to_anthropic_events`` — response-side translator
"""
from __future__ import annotations
import json
from collections.abc import AsyncGenerator
from typing import Any
import pytest
from routstr.upstream.gemini_messages import (
DUMMY_THOUGHT_SIGNATURE,
_openai_chunks_to_anthropic_events,
inject_thought_signatures,
)
# ---------------------------------------------------------------------------
# inject_thought_signatures
# ---------------------------------------------------------------------------
def test_inject_thought_signatures_adds_dummy_to_each_tool_call() -> None:
messages: list[dict[str, Any]] = [
{"role": "user", "content": "do thing"},
{
"role": "assistant",
"tool_calls": [
{
"id": "toolu_1",
"type": "function",
"function": {"name": "Bash", "arguments": "{}"},
},
{
"id": "toolu_2",
"type": "function",
"function": {"name": "Read", "arguments": "{}"},
},
],
},
]
inject_thought_signatures(messages)
for tc in messages[1]["tool_calls"]:
assert (
tc["extra_content"]["google"]["thought_signature"]
== DUMMY_THOUGHT_SIGNATURE
)
def test_inject_thought_signatures_preserves_existing_signature() -> None:
"""Don't clobber a real signature that came back from a prior turn."""
messages: list[dict[str, Any]] = [
{
"role": "assistant",
"tool_calls": [
{
"id": "tc1",
"type": "function",
"function": {"name": "fn", "arguments": "{}"},
"extra_content": {
"google": {"thought_signature": "real-signature"}
},
}
],
}
]
inject_thought_signatures(messages)
assert (
messages[0]["tool_calls"][0]["extra_content"]["google"][
"thought_signature"
]
== "real-signature"
)
def test_inject_thought_signatures_skips_messages_without_tool_calls() -> None:
messages: list[dict[str, Any]] = [
{"role": "user", "content": "hi"},
{"role": "assistant", "content": "hello"},
]
inject_thought_signatures(messages)
for m in messages:
assert "extra_content" not in m
def test_inject_thought_signatures_handles_malformed_extra_content() -> None:
"""If a caller already set ``extra_content`` to a non-dict (defensive),
we replace it instead of crashing."""
messages: list[dict[str, Any]] = [
{
"role": "assistant",
"tool_calls": [
{
"id": "tc",
"type": "function",
"function": {"name": "fn", "arguments": "{}"},
"extra_content": "garbage",
}
],
}
]
inject_thought_signatures(messages)
extra = messages[0]["tool_calls"][0]["extra_content"]
assert isinstance(extra, dict)
assert extra["google"]["thought_signature"] == DUMMY_THOUGHT_SIGNATURE
# ---------------------------------------------------------------------------
# _openai_chunks_to_anthropic_events
# ---------------------------------------------------------------------------
async def _lines(*chunks: dict | str) -> AsyncGenerator[str, None]:
"""Helper to wrap chunk dicts as SSE-style ``data:`` lines."""
for c in chunks:
if isinstance(c, dict):
yield f"data: {json.dumps(c)}"
else:
yield c
def _parse_anthropic_sse(blocks: list[bytes]) -> list[dict]:
"""Flatten a list of Anthropic SSE byte chunks into event dicts."""
events: list[dict] = []
for blob in blocks:
text = blob.decode()
for entry in text.split("\n\n"):
for line in entry.splitlines():
if line.startswith("data:"):
events.append(json.loads(line[5:].lstrip()))
return events
@pytest.mark.asyncio
async def test_translator_emits_text_only_response() -> None:
"""Plain text response: message_start → content_block_* (text) →
message_delta(end_turn) → message_stop."""
chunks: list[dict] = [
{
"id": "chatcmpl-1",
"model": "gemini-2.5-flash",
"choices": [{"index": 0, "delta": {"role": "assistant"}}],
},
{"choices": [{"delta": {"content": "Hello"}}]},
{"choices": [{"delta": {"content": ", world"}}]},
{
"choices": [{"delta": {}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 5, "completion_tokens": 7},
},
]
out = []
async for event_bytes in _openai_chunks_to_anthropic_events(
_lines(*chunks), requested_model="gemini-2.5-flash"
):
out.append(event_bytes)
events = _parse_anthropic_sse(out)
types = [e["type"] for e in events]
assert types == [
"message_start",
"content_block_start",
"content_block_delta",
"content_block_delta",
"content_block_stop",
"message_delta",
"message_stop",
]
# Text deltas concatenate to "Hello, world".
text_deltas = [
e["delta"]["text"]
for e in events
if e["type"] == "content_block_delta"
]
assert "".join(text_deltas) == "Hello, world"
# Stop reason was mapped from openai's "stop".
msg_delta = next(e for e in events if e["type"] == "message_delta")
assert msg_delta["delta"]["stop_reason"] == "end_turn"
assert msg_delta["usage"]["input_tokens"] == 5
assert msg_delta["usage"]["output_tokens"] == 7
@pytest.mark.asyncio
async def test_translator_emits_tool_use_block() -> None:
"""tool_calls split across deltas → tool_use content block with
accumulated input_json_delta and stop_reason='tool_use'."""
chunks: list[dict] = [
{
"id": "chatcmpl-2",
"model": "gemini-2.5-flash",
"choices": [{"delta": {"role": "assistant"}}],
},
{
"choices": [
{
"delta": {
"tool_calls": [
{
"index": 0,
"id": "call-abc",
"type": "function",
"function": {
"name": "Bash",
"arguments": '{"cmd":',
},
}
]
}
}
]
},
{
"choices": [
{
"delta": {
"tool_calls": [
{
"index": 0,
"function": {"arguments": ' "ls"}'},
}
]
}
}
]
},
{"choices": [{"delta": {}, "finish_reason": "tool_calls"}]},
]
out = []
async for event_bytes in _openai_chunks_to_anthropic_events(
_lines(*chunks), requested_model="gemini-2.5-flash"
):
out.append(event_bytes)
events = _parse_anthropic_sse(out)
types = [e["type"] for e in events]
assert types == [
"message_start",
"content_block_start",
"content_block_delta",
"content_block_delta",
"content_block_stop",
"message_delta",
"message_stop",
]
# Tool use block was opened with the right name.
cb_start = next(e for e in events if e["type"] == "content_block_start")
assert cb_start["content_block"]["type"] == "tool_use"
assert cb_start["content_block"]["name"] == "Bash"
assert cb_start["content_block"]["id"] == "call-abc"
# Argument deltas were forwarded as input_json_delta partials.
deltas = [e for e in events if e["type"] == "content_block_delta"]
assert all(d["delta"]["type"] == "input_json_delta" for d in deltas)
assert "".join(d["delta"]["partial_json"] for d in deltas) == (
'{"cmd": "ls"}'
)
# tool_calls finish_reason → tool_use stop_reason.
msg_delta = next(e for e in events if e["type"] == "message_delta")
assert msg_delta["delta"]["stop_reason"] == "tool_use"
@pytest.mark.asyncio
async def test_translator_handles_done_sentinel_and_blank_lines() -> None:
"""Spec edge cases from openai SSE: ``data: [DONE]``, blank lines,
invalid JSON. Translator should skip them gracefully."""
chunks: list[dict | str] = [
{
"id": "x",
"model": "m",
"choices": [{"delta": {"role": "assistant"}}],
},
{"choices": [{"delta": {"content": "ok"}}]},
"",
": comment",
"data: not-json",
"data: [DONE]",
{"choices": [{"delta": {}, "finish_reason": "stop"}]},
]
out = []
async for event_bytes in _openai_chunks_to_anthropic_events(
_lines(*chunks), requested_model=None
):
out.append(event_bytes)
events = _parse_anthropic_sse(out)
assert events[0]["type"] == "message_start"
assert events[-1]["type"] == "message_stop"
text = "".join(
e["delta"]["text"]
for e in events
if e["type"] == "content_block_delta"
)
assert text == "ok"

View File

@@ -74,3 +74,39 @@ async def test_get_balance_returns_none_on_connect_timeout(
balance = await provider.get_balance()
assert balance is None
def test_normalize_request_path_keeps_v1_prefix() -> None:
"""Routstr upstream stores ``base_url`` without ``/v1``; the prefix
must stay on the path so ``build_request_url`` produces ``/v1/<endpoint>``
instead of ``/<endpoint>`` (which the upstream Routstr 404s with HTML)."""
provider = RoutstrUpstreamProvider(
base_url="https://privateprovider.xyz", api_key="key"
)
assert provider.normalize_request_path("v1/messages") == "v1/messages"
assert provider.normalize_request_path("/v1/messages") == "v1/messages"
assert (
provider.normalize_request_path("v1/chat/completions")
== "v1/chat/completions"
)
def test_build_request_url_for_v1_messages() -> None:
"""Forwarding ``/v1/messages`` must hit the upstream's ``/v1/messages``."""
provider = RoutstrUpstreamProvider(
base_url="https://privateprovider.xyz", api_key="key"
)
normalized = provider.normalize_request_path("v1/messages")
assert (
provider.build_request_url(normalized)
== "https://privateprovider.xyz/v1/messages"
)
def test_supports_anthropic_messages_natively() -> None:
"""Routstr nodes serve ``/v1/messages`` directly, so the proxy must
forward as-is instead of round-tripping through litellm."""
assert RoutstrUpstreamProvider.supports_anthropic_messages is True

View File

@@ -123,6 +123,7 @@ async def test_swap_to_primary_mint_insufficient_for_fees() -> None:
mock_token_wallet = Mock()
mock_token_wallet.load_mint = AsyncMock()
mock_token_wallet.load_proofs = AsyncMock()
mock_token_wallet.get_fees_for_proofs = Mock(return_value=0)
mock_primary_wallet = Mock()
mock_primary_wallet.load_mint = AsyncMock()
@@ -166,6 +167,7 @@ async def test_swap_to_primary_mint_melt_error_wrapped() -> None:
mock_token_wallet = Mock()
mock_token_wallet.load_mint = AsyncMock()
mock_token_wallet.load_proofs = AsyncMock()
mock_token_wallet.get_fees_for_proofs = Mock(return_value=0)
mock_primary_wallet = Mock()
mock_primary_wallet.load_mint = AsyncMock()
@@ -220,6 +222,33 @@ async def test_recieve_token_untrusted_mint() -> None:
@pytest.mark.asyncio
@pytest.mark.asyncio
async def test_swap_to_primary_mint_already_on_primary() -> None:
from routstr.core.settings import settings
from routstr.wallet import swap_to_primary_mint
mock_token = Mock()
mock_token.mint = settings.primary_mint
mock_token.amount = 1000
mock_token.unit = "sat"
mock_token.proofs = []
mock_token_wallet = Mock()
mock_token_wallet.split = AsyncMock(return_value=None)
mock_token_wallet.request_mint = AsyncMock()
mock_token_wallet.melt_quote = AsyncMock()
with patch("routstr.wallet.get_wallet", AsyncMock(return_value=mock_token_wallet)):
amount, unit, mint = await swap_to_primary_mint(mock_token, mock_token_wallet)
assert amount == 1000
assert unit == "sat"
assert mint == settings.primary_mint
mock_token_wallet.split.assert_called_once()
mock_token_wallet.request_mint.assert_not_called()
mock_token_wallet.melt_quote.assert_not_called()
async def test_swap_to_primary_mint_success() -> None:
"""Test successful swap with dynamic fee calculation."""
from routstr.wallet import swap_to_primary_mint
@@ -234,6 +263,7 @@ async def test_swap_to_primary_mint_success() -> None:
mock_token_wallet = Mock()
mock_token_wallet.load_mint = AsyncMock()
mock_token_wallet.load_proofs = AsyncMock()
mock_token_wallet.get_fees_for_proofs = Mock(return_value=0)
mock_primary_wallet = Mock()
mock_primary_wallet.load_mint = AsyncMock()

View File

@@ -4,6 +4,7 @@ import * as React from 'react';
import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs';
import { ServerConfigSettings } from '@/components/settings/server-config-settings';
import { AdminSettings } from '@/components/settings/admin-settings';
import { CliTokensSettings } from '@/components/settings/cli-tokens-settings';
import { AppPageShell } from '@/components/app-page-shell';
import { PageHeader } from '@/components/page-header';
@@ -19,6 +20,7 @@ export default function SettingsPage() {
<TabsList variant='line' className='mb-4 w-full'>
<TabsTrigger value='admin'>Admin Settings</TabsTrigger>
<TabsTrigger value='server'>Server Config</TabsTrigger>
<TabsTrigger value='cli-tokens'>CLI Tokens</TabsTrigger>
</TabsList>
<TabsContent value='server'>
<ServerConfigSettings />
@@ -26,6 +28,9 @@ export default function SettingsPage() {
<TabsContent value='admin'>
<AdminSettings />
</TabsContent>
<TabsContent value='cli-tokens'>
<CliTokensSettings />
</TabsContent>
</Tabs>
</div>
</AppPageShell>

View File

@@ -1,7 +1,7 @@
'use client';
import { useState, useEffect } from 'react';
import { useQuery } from '@tanstack/react-query';
import { useQuery, keepPreviousData } from '@tanstack/react-query';
import { AppPageShell } from '@/components/app-page-shell';
import { PageHeader } from '@/components/page-header';
import {
@@ -22,6 +22,7 @@ import {
SelectValue,
} from '@/components/ui/select';
import { Badge } from '@/components/ui/badge';
import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs';
import {
Table,
TableBody,
@@ -30,7 +31,7 @@ import {
TableHeader,
TableRow,
} from '@/components/ui/table';
import { ScrollArea } from '@/components/ui/scroll-area';
import { ScrollArea, ScrollBar } from '@/components/ui/scroll-area';
import { Skeleton } from '@/components/ui/skeleton';
import {
Empty,
@@ -47,13 +48,328 @@ import {
Copy,
Check,
Receipt,
Key,
Zap,
ChevronLeft,
ChevronRight,
} from 'lucide-react';
import { AdminService, type Transaction } from '@/lib/api/services/admin';
import {
AdminService,
type Transaction,
type LightningInvoice,
} from '@/lib/api/services/admin';
import { format } from 'date-fns';
import { toast } from 'sonner';
const STORAGE_KEY = 'routstr-transaction-filters';
function TransactionTable({
transactions,
copiedId,
onCopy,
getStatusBadge,
}: {
transactions: Transaction[];
copiedId: string | null;
onCopy: (text: string, id: string) => void;
getStatusBadge: (tx: Transaction) => React.ReactNode;
}) {
if (transactions.length === 0) {
return (
<Empty className='py-8'>
<EmptyHeader>
<EmptyMedia variant='icon'>
<Receipt className='h-4 w-4' />
</EmptyMedia>
<EmptyTitle>No transactions found</EmptyTitle>
<EmptyDescription>
Try adjusting your filters or check back later.
</EmptyDescription>
</EmptyHeader>
</Empty>
);
}
return (
<ScrollArea className='h-[55svh] min-h-[420px] w-full sm:h-[600px]'>
<div className='min-w-[800px]'>
<Table>
<TableHeader>
<TableRow>
<TableHead>Type</TableHead>
<TableHead>Amount</TableHead>
<TableHead>Status</TableHead>
<TableHead>API Key</TableHead>
<TableHead>Request ID</TableHead>
<TableHead>Mint</TableHead>
<TableHead>Date</TableHead>
<TableHead className='text-right'>Actions</TableHead>
</TableRow>
</TableHeader>
<TableBody>
{transactions.map((tx) => (
<TableRow key={tx.id}>
<TableCell>
<div className='flex items-center gap-2'>
{tx.type === 'in' ? (
<ArrowDownLeft className='h-4 w-4 text-green-500' />
) : (
<ArrowUpRight className='h-4 w-4 text-blue-500' />
)}
<span className='capitalize'>{tx.type}</span>
</div>
</TableCell>
<TableCell className='font-mono'>
{tx.amount} {tx.unit}
</TableCell>
<TableCell>{getStatusBadge(tx)}</TableCell>
<TableCell>
{tx.api_key_hashed_key ? (
<div className='flex items-center gap-1 text-xs'>
<span className='max-w-[120px] truncate font-mono'>
{tx.api_key_hashed_key.slice(0, 12)}...
</span>
<Button
variant='ghost'
size='icon'
className='h-4 w-4'
onClick={() =>
onCopy(tx.api_key_hashed_key!, tx.id + '-apikey')
}
>
{copiedId === tx.id + '-apikey' ? (
<Check className='h-3 w-3' />
) : (
<Copy className='h-3 w-3' />
)}
</Button>
</div>
) : (
<span className='text-muted-foreground text-xs'></span>
)}
</TableCell>
<TableCell>
{tx.request_id ? (
<div className='flex items-center gap-1 text-xs'>
<span className='max-w-[150px] truncate font-mono'>
{tx.request_id}
</span>
<Button
variant='ghost'
size='icon'
className='h-4 w-4'
onClick={() => onCopy(tx.request_id!, tx.id + '-req')}
>
{copiedId === tx.id + '-req' ? (
<Check className='h-3 w-3' />
) : (
<Copy className='h-3 w-3' />
)}
</Button>
</div>
) : (
<span className='text-muted-foreground text-xs'></span>
)}
</TableCell>
<TableCell>
<div className='flex max-w-[150px] items-center gap-1 truncate text-xs'>
<span className='truncate'>{tx.mint_url}</span>
</div>
</TableCell>
<TableCell className='text-xs whitespace-nowrap'>
{format(tx.created_at * 1000, 'yyyy-MM-dd HH:mm:ss')}
</TableCell>
<TableCell className='text-right'>
<Button
variant='ghost'
size='icon'
className='h-8 w-8'
onClick={() => onCopy(tx.token, tx.id + '-token')}
title='Copy Token'
>
{copiedId === tx.id + '-token' ? (
<Check className='h-4 w-4' />
) : (
<Copy className='h-4 w-4' />
)}
</Button>
</TableCell>
</TableRow>
))}
</TableBody>
</Table>
</div>
<ScrollBar orientation='horizontal' />
</ScrollArea>
);
}
function LightningInvoiceTable({
invoices,
copiedId,
onCopy,
}: {
invoices: LightningInvoice[];
copiedId: string | null;
onCopy: (text: string, id: string) => void;
}) {
if (invoices.length === 0) {
return (
<Empty className='py-8'>
<EmptyHeader>
<EmptyMedia variant='icon'>
<Zap className='h-4 w-4' />
</EmptyMedia>
<EmptyTitle>No invoices found</EmptyTitle>
<EmptyDescription>
Lightning invoices created via /lightning/invoice will show here.
</EmptyDescription>
</EmptyHeader>
</Empty>
);
}
const statusBadge = (status: LightningInvoice['status']) => {
if (status === 'paid')
return (
<Badge
variant='outline'
className='border-green-500/20 bg-green-500/10 text-green-500'
>
Paid
</Badge>
);
if (status === 'expired')
return (
<Badge
variant='outline'
className='border-red-500/20 bg-red-500/10 text-red-500'
>
Expired
</Badge>
);
if (status === 'cancelled')
return (
<Badge
variant='outline'
className='border-gray-500/20 bg-gray-500/10 text-gray-500'
>
Cancelled
</Badge>
);
return (
<Badge
variant='outline'
className='border-blue-500/20 bg-blue-500/10 text-blue-500'
>
Pending
</Badge>
);
};
return (
<ScrollArea className='h-[55svh] min-h-[420px] w-full sm:h-[600px]'>
<div className='min-w-[900px]'>
<Table>
<TableHeader>
<TableRow>
<TableHead>Purpose</TableHead>
<TableHead>Amount</TableHead>
<TableHead>Status</TableHead>
<TableHead>API Key</TableHead>
<TableHead>Payment Hash</TableHead>
<TableHead>Created</TableHead>
<TableHead>Paid</TableHead>
<TableHead className='text-right'>Actions</TableHead>
</TableRow>
</TableHeader>
<TableBody>
{invoices.map((inv) => (
<TableRow key={inv.id}>
<TableCell>
<span className='capitalize'>{inv.purpose}</span>
</TableCell>
<TableCell className='font-mono'>
{inv.amount_sats} sat
</TableCell>
<TableCell>{statusBadge(inv.status)}</TableCell>
<TableCell>
{inv.api_key_hash ? (
<div className='flex items-center gap-1 text-xs'>
<span className='max-w-[120px] truncate font-mono'>
{inv.api_key_hash.slice(0, 12)}...
</span>
<Button
variant='ghost'
size='icon'
className='h-4 w-4'
onClick={() =>
onCopy(inv.api_key_hash!, inv.id + '-apikey')
}
>
{copiedId === inv.id + '-apikey' ? (
<Check className='h-3 w-3' />
) : (
<Copy className='h-3 w-3' />
)}
</Button>
</div>
) : (
<span className='text-muted-foreground text-xs'></span>
)}
</TableCell>
<TableCell>
<div className='flex items-center gap-1 text-xs'>
<span className='max-w-[140px] truncate font-mono'>
{inv.payment_hash.slice(0, 14)}...
</span>
<Button
variant='ghost'
size='icon'
className='h-4 w-4'
onClick={() => onCopy(inv.payment_hash, inv.id + '-hash')}
>
{copiedId === inv.id + '-hash' ? (
<Check className='h-3 w-3' />
) : (
<Copy className='h-3 w-3' />
)}
</Button>
</div>
</TableCell>
<TableCell className='text-xs whitespace-nowrap'>
{format(inv.created_at * 1000, 'yyyy-MM-dd HH:mm:ss')}
</TableCell>
<TableCell className='text-xs whitespace-nowrap'>
{inv.paid_at
? format(inv.paid_at * 1000, 'yyyy-MM-dd HH:mm:ss')
: '—'}
</TableCell>
<TableCell className='text-right'>
<Button
variant='ghost'
size='icon'
className='h-8 w-8'
onClick={() => onCopy(inv.bolt11, inv.id + '-bolt11')}
title='Copy BOLT11'
>
{copiedId === inv.id + '-bolt11' ? (
<Check className='h-4 w-4' />
) : (
<Copy className='h-4 w-4' />
)}
</Button>
</TableCell>
</TableRow>
))}
</TableBody>
</Table>
</div>
<ScrollBar orientation='horizontal' />
</ScrollArea>
);
}
export default function TransactionsPage() {
const [search, setSearch] = useState('');
const [type, setType] = useState<string>('all');
@@ -81,21 +397,89 @@ export default function TransactionsPage() {
localStorage.setItem(STORAGE_KEY, JSON.stringify(filters));
}, [search, type, status]);
const { data, isLoading, refetch, isRefetching } = useQuery({
queryKey: ['transactions', type, status, search],
const PAGE_SIZE = 50;
const [activeTab, setActiveTab] = useState<string>('x-cashu');
const [xcashuPage, setXcashuPage] = useState(0);
const [apikeyPage, setApikeyPage] = useState(0);
const [lightningPage, setLightningPage] = useState(0);
const typeParam = type === 'all' ? undefined : type;
const statusParam = status === 'all' ? undefined : status;
const searchParam = search || undefined;
const xcashuQuery = useQuery({
queryKey: [
'transactions',
'x-cashu',
typeParam,
statusParam,
searchParam,
xcashuPage,
],
queryFn: () =>
AdminService.getTransactions(
type === 'all' ? undefined : type,
status === 'all' ? undefined : status,
search || undefined,
100
typeParam,
statusParam,
searchParam,
'x-cashu',
PAGE_SIZE,
xcashuPage * PAGE_SIZE
),
placeholderData: keepPreviousData,
});
const apikeyQuery = useQuery({
queryKey: [
'transactions',
'apikey',
typeParam,
statusParam,
searchParam,
apikeyPage,
],
queryFn: () =>
AdminService.getTransactions(
typeParam,
statusParam,
searchParam,
'apikey',
PAGE_SIZE,
apikeyPage * PAGE_SIZE
),
placeholderData: keepPreviousData,
});
const LIGHTNING_STATUSES = ['pending', 'paid', 'expired', 'cancelled'];
const lightningStatusParam = LIGHTNING_STATUSES.includes(status)
? status
: undefined;
const lightningQuery = useQuery({
queryKey: [
'lightning-invoices',
lightningStatusParam,
searchParam,
lightningPage,
],
queryFn: () =>
AdminService.getLightningInvoices(
lightningStatusParam,
undefined,
searchParam,
PAGE_SIZE,
lightningPage * PAGE_SIZE
),
placeholderData: keepPreviousData,
refetchInterval: 10000,
});
const handleClearFilters = () => {
setSearch('');
setType('all');
setStatus('all');
setXcashuPage(0);
setApikeyPage(0);
setLightningPage(0);
};
const copyToClipboard = (text: string, id: string) => {
@@ -145,15 +529,96 @@ export default function TransactionsPage() {
.filter(Boolean)
.join(' • ');
// Reset pages when filters change
useEffect(() => {
setXcashuPage(0);
setApikeyPage(0);
setLightningPage(0);
}, [type, status, search]);
const isRefetching =
xcashuQuery.isRefetching ||
apikeyQuery.isRefetching ||
lightningQuery.isRefetching;
const renderCardContent = (
query: typeof xcashuQuery,
page: number,
setPage: (p: number) => void
) => {
if (query.isLoading) {
return (
<div className='space-y-2'>
{Array.from({ length: 8 }).map((_, index) => (
<Skeleton
key={`tx-loading-${index}`}
className='h-16 w-full rounded-lg'
/>
))}
</div>
);
}
const transactions = query.data?.transactions ?? [];
const total = query.data?.total ?? 0;
const totalPages = Math.ceil(total / PAGE_SIZE);
return (
<>
{totalPages > 1 && (
<div className='flex flex-col gap-2 border-b pb-3 sm:flex-row sm:items-center sm:justify-between'>
<span className='text-muted-foreground text-xs sm:text-sm'>
{page * PAGE_SIZE + 1}{Math.min((page + 1) * PAGE_SIZE, total)}{' '}
of {total}
</span>
<div className='flex items-center gap-2'>
<Button
variant='outline'
size='sm'
disabled={page === 0}
onClick={() => setPage(page - 1)}
>
<ChevronLeft className='h-4 w-4' />
<span className='hidden sm:inline'>Previous</span>
</Button>
<span className='text-xs sm:text-sm'>
{page + 1} / {totalPages}
</span>
<Button
variant='outline'
size='sm'
disabled={page >= totalPages - 1}
onClick={() => setPage(page + 1)}
>
<span className='hidden sm:inline'>Next</span>
<ChevronRight className='h-4 w-4' />
</Button>
</div>
</div>
)}
<TransactionTable
transactions={transactions}
copiedId={copiedId}
onCopy={copyToClipboard}
getStatusBadge={getStatusBadge}
/>
</>
);
};
return (
<AppPageShell contentClassName='mx-auto w-full max-w-5xl overflow-x-hidden'>
<div className='space-y-6'>
<PageHeader
title='X-Cashu Transactions'
description='View all incoming and outgoing X-Cashu token transactions.'
title='Cashu Transactions'
description='View all incoming and outgoing Cashu token transactions.'
actions={
<Button
onClick={() => refetch()}
onClick={() => {
xcashuQuery.refetch();
apikeyQuery.refetch();
lightningQuery.refetch();
}}
variant='outline'
size='sm'
disabled={isRefetching}
@@ -181,7 +646,7 @@ export default function TransactionsPage() {
<Search className='text-muted-foreground absolute top-2.5 left-2.5 h-4 w-4' />
<Input
id='search'
placeholder='Search by ID, token or request ID...'
placeholder='Search by ID, token, request ID or key hash...'
className='pl-8'
value={search}
onChange={(e) => setSearch(e.target.value)}
@@ -212,6 +677,11 @@ export default function TransactionsPage() {
<SelectItem value='pending'>Pending</SelectItem>
<SelectItem value='collected'>Collected</SelectItem>
<SelectItem value='swept'>Swept</SelectItem>
<SelectItem value='paid'>Paid (Lightning)</SelectItem>
<SelectItem value='expired'>Expired (Lightning)</SelectItem>
<SelectItem value='cancelled'>
Cancelled (Lightning)
</SelectItem>
</SelectContent>
</Select>
</div>
@@ -228,138 +698,152 @@ export default function TransactionsPage() {
</CardContent>
</Card>
<Card>
<CardHeader>
<div className='flex flex-col items-start gap-2 sm:flex-row sm:items-center sm:justify-between'>
<CardTitle>Transaction History</CardTitle>
{data && (
<Badge variant='secondary'>
{data.transactions.length} entries
<Tabs
defaultValue='x-cashu'
value={activeTab}
onValueChange={setActiveTab}
>
<TabsList className='mb-4'>
<TabsTrigger value='x-cashu' className='flex items-center gap-2'>
<Zap className='h-4 w-4' />
X-Cashu
{xcashuQuery.data && (
<Badge variant='secondary' className='ml-1'>
{xcashuQuery.data.total}
</Badge>
)}
</div>
{hasActiveFilters && (
<CardDescription>
Showing transactions filtered by {activeFilterDescription}
</CardDescription>
)}
</CardHeader>
<CardContent className='overflow-hidden'>
{isLoading ? (
<div className='space-y-2'>
{Array.from({ length: 8 }).map((_, index) => (
<Skeleton
key={`tx-loading-${index}`}
className='h-16 w-full rounded-lg'
/>
))}
</div>
) : data?.transactions && data.transactions.length > 0 ? (
<ScrollArea className='h-[55svh] min-h-[420px] w-full sm:h-[600px]'>
<Table>
<TableHeader>
<TableRow>
<TableHead>Type</TableHead>
<TableHead>Amount</TableHead>
<TableHead>Status</TableHead>
<TableHead>Request ID</TableHead>
<TableHead>Mint</TableHead>
<TableHead>Date</TableHead>
<TableHead className='text-right'>Actions</TableHead>
</TableRow>
</TableHeader>
<TableBody>
{data.transactions.map((tx) => (
<TableRow key={tx.id}>
<TableCell>
<div className='flex items-center gap-2'>
{tx.type === 'in' ? (
<ArrowDownLeft className='h-4 w-4 text-green-500' />
) : (
<ArrowUpRight className='h-4 w-4 text-blue-500' />
)}
<span className='capitalize'>{tx.type}</span>
</div>
</TableCell>
<TableCell className='font-mono'>
{tx.amount} {tx.unit}
</TableCell>
<TableCell>{getStatusBadge(tx)}</TableCell>
<TableCell>
{tx.request_id ? (
<div className='flex items-center gap-1 text-xs'>
<span className='max-w-[150px] truncate font-mono'>
{tx.request_id}
</span>
<Button
variant='ghost'
size='icon'
className='h-4 w-4'
onClick={() =>
copyToClipboard(
tx.request_id!,
tx.id + '-req'
)
}
>
{copiedId === tx.id + '-req' ? (
<Check className='h-3 w-3' />
) : (
<Copy className='h-3 w-3' />
)}
</Button>
</div>
) : (
<span className='text-muted-foreground text-xs'>
</span>
)}
</TableCell>
<TableCell>
<div className='flex max-w-[150px] items-center gap-1 truncate text-xs'>
<span className='truncate'>{tx.mint_url}</span>
</div>
</TableCell>
<TableCell className='text-xs whitespace-nowrap'>
{format(tx.created_at * 1000, 'yyyy-MM-dd HH:mm:ss')}
</TableCell>
<TableCell className='text-right'>
<Button
variant='ghost'
size='icon'
className='h-8 w-8'
onClick={() =>
copyToClipboard(tx.token, tx.id + '-token')
}
title='Copy Token'
>
{copiedId === tx.id + '-token' ? (
<Check className='h-4 w-4' />
) : (
<Copy className='h-4 w-4' />
)}
</Button>
</TableCell>
</TableRow>
</TabsTrigger>
<TabsTrigger value='apikey' className='flex items-center gap-2'>
<Key className='h-4 w-4' />
API Key
{apikeyQuery.data && (
<Badge variant='secondary' className='ml-1'>
{apikeyQuery.data.total}
</Badge>
)}
</TabsTrigger>
<TabsTrigger value='lightning' className='flex items-center gap-2'>
<Zap className='h-4 w-4' />
Lightning
{lightningQuery.data && (
<Badge variant='secondary' className='ml-1'>
{lightningQuery.data.total}
</Badge>
)}
</TabsTrigger>
</TabsList>
<TabsContent value='x-cashu'>
<Card>
<CardHeader>
<div className='flex flex-col items-start gap-2 sm:flex-row sm:items-center sm:justify-between'>
<CardTitle>X-Cashu Transaction History</CardTitle>
{hasActiveFilters && (
<CardDescription>
Filtered by {activeFilterDescription}
</CardDescription>
)}
</div>
</CardHeader>
<CardContent className='overflow-hidden'>
{renderCardContent(xcashuQuery, xcashuPage, setXcashuPage)}
</CardContent>
</Card>
</TabsContent>
<TabsContent value='apikey'>
<Card>
<CardHeader>
<div className='flex flex-col items-start gap-2 sm:flex-row sm:items-center sm:justify-between'>
<CardTitle>API Key Transaction History</CardTitle>
{hasActiveFilters && (
<CardDescription>
Filtered by {activeFilterDescription}
</CardDescription>
)}
</div>
</CardHeader>
<CardContent className='overflow-hidden'>
{renderCardContent(apikeyQuery, apikeyPage, setApikeyPage)}
</CardContent>
</Card>
</TabsContent>
<TabsContent value='lightning'>
<Card>
<CardHeader>
<div className='flex flex-col items-start gap-2 sm:flex-row sm:items-center sm:justify-between'>
<CardTitle>Lightning Invoice History</CardTitle>
<CardDescription>
Auto-refreshing every 10s. Paid invoices credit balance
automatically.
</CardDescription>
</div>
</CardHeader>
<CardContent className='overflow-hidden'>
{lightningQuery.isLoading ? (
<div className='space-y-2'>
{Array.from({ length: 8 }).map((_, index) => (
<Skeleton
key={`ln-loading-${index}`}
className='h-16 w-full rounded-lg'
/>
))}
</TableBody>
</Table>
</ScrollArea>
) : (
<Empty className='py-8'>
<EmptyHeader>
<EmptyMedia variant='icon'>
<Receipt className='h-4 w-4' />
</EmptyMedia>
<EmptyTitle>No transactions found</EmptyTitle>
<EmptyDescription>
Try adjusting your filters or check back later.
</EmptyDescription>
</EmptyHeader>
</Empty>
)}
</CardContent>
</Card>
</div>
) : (
<>
{(() => {
const total = lightningQuery.data?.total ?? 0;
const totalPages = Math.ceil(total / PAGE_SIZE);
if (totalPages <= 1) return null;
return (
<div className='flex flex-col gap-2 border-b pb-3 sm:flex-row sm:items-center sm:justify-between'>
<span className='text-muted-foreground text-xs sm:text-sm'>
{lightningPage * PAGE_SIZE + 1}
{Math.min((lightningPage + 1) * PAGE_SIZE, total)}{' '}
of {total}
</span>
<div className='flex items-center gap-2'>
<Button
variant='outline'
size='sm'
disabled={lightningPage === 0}
onClick={() =>
setLightningPage(lightningPage - 1)
}
>
<ChevronLeft className='h-4 w-4' />
<span className='hidden sm:inline'>Previous</span>
</Button>
<span className='text-xs sm:text-sm'>
{lightningPage + 1} / {totalPages}
</span>
<Button
variant='outline'
size='sm'
disabled={lightningPage >= totalPages - 1}
onClick={() =>
setLightningPage(lightningPage + 1)
}
>
<span className='hidden sm:inline'>Next</span>
<ChevronRight className='h-4 w-4' />
</Button>
</div>
</div>
);
})()}
<LightningInvoiceTable
invoices={lightningQuery.data?.invoices ?? []}
copiedId={copiedId}
onCopy={copyToClipboard}
/>
</>
)}
</CardContent>
</Card>
</TabsContent>
</Tabs>
</div>
</AppPageShell>
);

View File

@@ -39,7 +39,7 @@ import {
FormMessage,
} from '@/components/ui/form';
import { Switch } from '@/components/ui/switch';
import { Loader2, Plus } from 'lucide-react';
import { Check, Copy, Loader2, Plus } from 'lucide-react';
import { toast } from 'sonner';
import { AdminService, type AdminModel } from '@/lib/api/services/admin';
@@ -64,6 +64,7 @@ const FormSchema = z.object({
instruct_type: z.string().default(''),
canonical_slug: z.string().default(''),
alias_ids_raw: z.string().default(''),
forwarded_model_id: z.string().default(''),
upstream_provider_id: z.string().default(''),
input_cost: z.coerce.number().min(0).default(0),
output_cost: z.coerce.number().min(0).default(0),
@@ -104,6 +105,7 @@ export function AddProviderModelDialog({
const [isPresetOpen, setIsPresetOpen] = useState(false);
const [selectedPresetLabel, setSelectedPresetLabel] =
useState('Select a preset');
const [forwardedModelIdCopied, setForwardedModelIdCopied] = useState(false);
const form = useForm<FormData>({
resolver: zodResolver(FormSchema) as never,
@@ -119,6 +121,7 @@ export function AddProviderModelDialog({
instruct_type: '',
canonical_slug: '',
alias_ids_raw: '',
forwarded_model_id: '',
upstream_provider_id: '',
input_cost: 0,
output_cost: 0,
@@ -180,6 +183,7 @@ export function AddProviderModelDialog({
: '',
canonical_slug: initialData.canonical_slug || '',
alias_ids_raw: listToString(initialData.alias_ids),
forwarded_model_id: initialData.forwarded_model_id || initialData.id,
upstream_provider_id:
typeof initialData.upstream_provider_id === 'string'
? initialData.upstream_provider_id
@@ -223,6 +227,7 @@ export function AddProviderModelDialog({
instruct_type: '',
canonical_slug: '',
alias_ids_raw: '',
forwarded_model_id: '',
upstream_provider_id: '',
input_cost: 0,
output_cost: 0,
@@ -280,6 +285,7 @@ export function AddProviderModelDialog({
);
form.setValue('canonical_slug', model.canonical_slug || '');
form.setValue('alias_ids_raw', listToString(model.alias_ids));
form.setValue('forwarded_model_id', model.forwarded_model_id || model.id);
form.setValue(
'upstream_provider_id',
typeof model.upstream_provider_id === 'string'
@@ -385,6 +391,7 @@ export function AddProviderModelDialog({
canonical_slug: data.canonical_slug?.trim() || null,
alias_ids: listFromString(data.alias_ids_raw || ''),
enabled: data.enabled,
forwarded_model_id: data.forwarded_model_id?.trim() || data.id,
};
if (isEdit) {
@@ -520,6 +527,53 @@ export function AddProviderModelDialog({
)}
/>
<FormField
control={form.control}
name='forwarded_model_id'
render={({ field }) => {
const handleCopy = () => {
const value = field.value || form.getValues('id');
if (!value) return;
navigator.clipboard.writeText(value);
setForwardedModelIdCopied(true);
setTimeout(() => setForwardedModelIdCopied(false), 1500);
};
return (
<FormItem>
<FormLabel>Client Alias ID</FormLabel>
<FormControl>
<div className='flex gap-2'>
<Input
placeholder={
form.watch('id') || 'e.g., openai/gpt-4o'
}
{...field}
/>
<Button
type='button'
variant='outline'
size='icon'
onClick={handleCopy}
title='Copy model ID'
>
{forwardedModelIdCopied ? (
<Check className='h-4 w-4 text-green-500' />
) : (
<Copy className='h-4 w-4' />
)}
</Button>
</div>
</FormControl>
<FormDescription>
Alternate ID that clients can use to reference this
model. Defaults to the model&apos;s own ID.
</FormDescription>
<FormMessage />
</FormItem>
);
}}
/>
<FormField
control={form.control}
name='name'

View File

@@ -471,7 +471,11 @@ export function ApiEndpointTester({ models }: ApiEndpointTesterProps) {
testEndpointMutation.mutate(requestData);
};
const enabledModels = models.filter((model) => model.isEnabled);
const enabledModels = Array.from(
new Map(
models.filter((model) => model.isEnabled).map((m) => [m.id, m])
).values()
);
const credentials = selectedModel ? getModelCredentials(selectedModel) : null;
const endpointUrl = credentials
? buildEndpointUrl(

View File

@@ -21,6 +21,7 @@ import { adminLogout } from '@/lib/api/services/auth';
import { Button } from '@/components/ui/button';
import { CurrencyToggle } from '@/components/currency-toggle';
import { ThemeToggle } from '@/components/theme-toggle';
import { VersionStatus } from '@/components/version-status';
import {
Sheet,
SheetClose,
@@ -90,8 +91,18 @@ export function AppPageShell({
isSidebarCollapsed && 'px-0'
)}
>
<div className='flex items-center gap-2'>
<div className='flex min-w-0 flex-1 items-center gap-2 overflow-hidden'>
<div
className={cn(
'flex items-center gap-2',
isSidebarCollapsed && 'justify-center'
)}
>
<div
className={cn(
'flex min-w-0 items-center gap-2 overflow-hidden',
!isSidebarCollapsed && 'flex-1'
)}
>
<Image
src='/icon.ico'
alt='Routstr Node'
@@ -110,26 +121,9 @@ export function AppPageShell({
<h1 className='truncate text-lg font-semibold tracking-tight whitespace-nowrap'>
Routstr Node
</h1>
<VersionStatus className='mt-0.5' />
</div>
</div>
<Button
variant='ghost'
size='icon'
className={cn(
'text-muted-foreground hover:text-foreground h-8 w-8 shrink-0 transition-transform duration-300 ease-in-out',
isSidebarCollapsed ? 'mx-auto' : '-mr-1 ml-auto'
)}
onClick={() => setIsSidebarCollapsed((current) => !current)}
>
{isSidebarCollapsed ? (
<PanelLeftOpenIcon className='h-4 w-4' />
) : (
<PanelLeftCloseIcon className='h-4 w-4' />
)}
<span className='sr-only'>
{isSidebarCollapsed ? 'Expand sidebar' : 'Collapse sidebar'}
</span>
</Button>
</div>
</div>
@@ -230,6 +224,28 @@ export function AppPageShell({
</Button>
</div>
)}
<Button
variant='ghost'
size={isSidebarCollapsed ? 'icon' : 'sm'}
className={cn(
'text-muted-foreground hover:text-foreground transition-[width,padding] duration-300 ease-in-out',
isSidebarCollapsed
? 'mx-auto h-8 w-8'
: 'h-8 w-full justify-start gap-1.5 rounded-md px-2.5 text-[11px]'
)}
onClick={() => setIsSidebarCollapsed((current) => !current)}
>
{isSidebarCollapsed ? (
<PanelLeftOpenIcon className='h-4 w-4' />
) : (
<PanelLeftCloseIcon className='h-4 w-4' />
)}
{isSidebarCollapsed ? (
<span className='sr-only'>Expand sidebar</span>
) : (
'Collapse'
)}
</Button>
</div>
</aside>
@@ -279,9 +295,12 @@ export function AppPageShell({
height={24}
className='rounded-sm'
/>
<p className='truncate text-base font-medium tracking-tight'>
Routstr Node
</p>
<div className='min-w-0'>
<p className='truncate text-base font-medium tracking-tight'>
Routstr Node
</p>
<VersionStatus className='mt-0.5' />
</div>
</div>
<SheetClose asChild>
<Button

View File

@@ -1,577 +0,0 @@
'use client';
import React, { useState, useEffect, useCallback } from 'react';
import { useForm } from 'react-hook-form';
import { zodResolver } from '@hookform/resolvers/zod';
import { z } from 'zod';
import { type Model } from '@/lib/api/schemas/models';
import { AdminService, type AdminModel } from '@/lib/api/services/admin';
import { Button } from '@/components/ui/button';
import { Input } from '@/components/ui/input';
import { Textarea } from '@/components/ui/textarea';
import {
Dialog,
DialogContent,
DialogDescription,
DialogHeader,
DialogTitle,
} from '@/components/ui/dialog';
import {
Form,
FormControl,
FormDescription,
FormField,
FormItem,
FormLabel,
FormMessage,
} from '@/components/ui/form';
import { Edit3, Loader2 } from 'lucide-react';
import { toast } from 'sonner';
import { Switch } from '@/components/ui/switch';
const EditModelFormSchema = z.object({
name: z.string().min(1, 'Name is required'),
description: z.string().optional(),
context_length: z.number().min(0),
prompt: z.number().min(0),
completion: z.number().min(0),
enabled: z.boolean(),
});
type EditModelFormData = z.infer<typeof EditModelFormSchema>;
const roundToFiveDecimals = (value: number | undefined | null): number => {
if (value === undefined || value === null || isNaN(value)) {
return 0;
}
return Math.round(value * 100000) / 100000;
};
const toNumber = (value: unknown, fallback = 0): number => {
if (typeof value === 'number' && Number.isFinite(value)) {
return value;
}
if (typeof value === 'string') {
const parsed = Number(value);
if (Number.isFinite(parsed)) {
return parsed;
}
}
return fallback;
};
const toStringArray = (value: unknown, fallback: string[]): string[] => {
if (!Array.isArray(value)) {
return fallback;
}
const filtered = value.filter(
(item): item is string => typeof item === 'string'
);
return filtered.length > 0 ? filtered : fallback;
};
interface EditModelFormProps {
model: Model;
providerId?: number;
onModelUpdate?: () => void;
onCancel?: () => void;
isOpen: boolean;
}
interface AdminModelData {
id: string;
name: string;
description?: string;
created: number;
context_length: number;
architecture: {
modality: string;
input_modalities: string[];
output_modalities: string[];
tokenizer: string;
instruct_type: string | null;
};
pricing: {
prompt: number;
completion: number;
request: number;
image: number;
web_search: number;
internal_reasoning: number;
};
per_request_limits: null | undefined;
top_provider: null | undefined;
upstream_provider_id: number;
enabled: boolean;
}
const normalizeAdminModelData = (
adminModel: AdminModel,
fallbackModel: Model,
providerId: number
): AdminModelData => {
const pricingRecord =
adminModel.pricing && typeof adminModel.pricing === 'object'
? (adminModel.pricing as Record<string, unknown>)
: {};
const architectureRecord =
adminModel.architecture && typeof adminModel.architecture === 'object'
? (adminModel.architecture as Record<string, unknown>)
: {};
return {
id: adminModel.id,
name: adminModel.name,
description: adminModel.description || '',
created: toNumber(adminModel.created, Math.floor(Date.now() / 1000)),
context_length: Math.max(
0,
Math.trunc(
toNumber(adminModel.context_length, fallbackModel.contextLength || 4096)
)
),
architecture: {
modality:
typeof architectureRecord.modality === 'string'
? architectureRecord.modality
: fallbackModel.modelType || 'text',
input_modalities: toStringArray(architectureRecord.input_modalities, [
fallbackModel.modelType || 'text',
]),
output_modalities: toStringArray(architectureRecord.output_modalities, [
fallbackModel.modelType || 'text',
]),
tokenizer:
typeof architectureRecord.tokenizer === 'string'
? architectureRecord.tokenizer
: '',
instruct_type:
typeof architectureRecord.instruct_type === 'string'
? architectureRecord.instruct_type
: null,
},
pricing: {
prompt: roundToFiveDecimals(
toNumber(pricingRecord.prompt, fallbackModel.input_cost)
),
completion: roundToFiveDecimals(
toNumber(pricingRecord.completion, fallbackModel.output_cost)
),
request: toNumber(pricingRecord.request, 0),
image: toNumber(pricingRecord.image, 0),
web_search: toNumber(pricingRecord.web_search, 0),
internal_reasoning: toNumber(pricingRecord.internal_reasoning, 0),
},
per_request_limits:
adminModel.per_request_limits === null ||
adminModel.per_request_limits === undefined
? adminModel.per_request_limits
: null,
top_provider:
adminModel.top_provider === null || adminModel.top_provider === undefined
? adminModel.top_provider
: null,
upstream_provider_id:
typeof adminModel.upstream_provider_id === 'number'
? adminModel.upstream_provider_id
: providerId,
enabled: adminModel.enabled !== false,
};
};
export function EditModelForm({
model,
providerId,
onModelUpdate,
onCancel,
isOpen,
}: EditModelFormProps) {
const [isSubmitting, setIsSubmitting] = useState(false);
const [adminModelData, setAdminModelData] = useState<AdminModelData | null>(
null
);
const [isNewOverride, setIsNewOverride] = useState(false);
const form = useForm<EditModelFormData>({
resolver: zodResolver(EditModelFormSchema),
defaultValues: {
name: model.name,
description: model.description || '',
context_length: model.contextLength || 4096,
prompt: roundToFiveDecimals(model.input_cost),
completion: roundToFiveDecimals(model.output_cost),
enabled: model.isEnabled !== false,
},
});
const loadAdminModel = useCallback(async () => {
if (!providerId) {
console.error('loadAdminModel called without providerId');
return;
}
try {
const adminModel = await AdminService.getProviderModel(
providerId,
model.id
);
const normalizedAdminModel = normalizeAdminModelData(
adminModel,
model,
providerId
);
setAdminModelData(normalizedAdminModel);
setIsNewOverride(false);
form.reset({
name: normalizedAdminModel.name,
description: normalizedAdminModel.description || '',
context_length: normalizedAdminModel.context_length,
prompt: normalizedAdminModel.pricing.prompt,
completion: normalizedAdminModel.pricing.completion,
enabled: normalizedAdminModel.enabled !== false,
});
} catch {
setIsNewOverride(true);
setAdminModelData({
id: model.full_name,
name: model.name,
description: model.description || '',
created: Math.floor(Date.now() / 1000),
context_length: model.contextLength || 4096,
architecture: {
modality: model.modelType || 'text',
input_modalities: [model.modelType || 'text'],
output_modalities: [model.modelType || 'text'],
tokenizer: '',
instruct_type: null,
},
pricing: {
prompt: roundToFiveDecimals(model.input_cost),
completion: roundToFiveDecimals(model.output_cost),
request: 0,
image: 0,
web_search: 0,
internal_reasoning: 0,
},
per_request_limits: null,
top_provider: null,
upstream_provider_id: providerId,
enabled: model.isEnabled !== false,
});
form.reset({
name: model.name,
description: model.description || '',
context_length: model.contextLength || 4096,
prompt: roundToFiveDecimals(model.input_cost),
completion: roundToFiveDecimals(model.output_cost),
enabled: model.isEnabled !== false,
});
}
}, [providerId, model, form]);
useEffect(() => {
if (isOpen && providerId) {
loadAdminModel();
} else if (isOpen && !providerId) {
console.error('EditModelForm opened without providerId', {
model,
providerId,
});
toast.error('Missing provider information for this model');
}
}, [isOpen, providerId, model, loadAdminModel]);
const onSubmit = async (data: EditModelFormData) => {
if (!providerId) {
console.error('onSubmit called without providerId', {
model,
providerId,
});
toast.error('Missing provider ID - cannot update model');
return;
}
if (!adminModelData) {
console.error('onSubmit called without adminModelData', {
model,
providerId,
adminModelData,
});
toast.error('Model data not loaded - please try reopening the form');
return;
}
setIsSubmitting(true);
try {
const payload = {
id: adminModelData.id,
name: data.name,
description: data.description || '',
created: adminModelData.created || Math.floor(Date.now() / 1000),
context_length: data.context_length,
architecture: adminModelData.architecture || {
modality: 'text',
input_modalities: ['text'],
output_modalities: ['text'],
tokenizer: '',
instruct_type: null,
},
pricing: {
prompt: roundToFiveDecimals(data.prompt),
completion: roundToFiveDecimals(data.completion),
request: 0,
image: 0,
web_search: 0,
internal_reasoning: 0,
},
per_request_limits: adminModelData.per_request_limits,
top_provider: adminModelData.top_provider,
upstream_provider_id: providerId,
enabled: data.enabled,
};
if (isNewOverride) {
await AdminService.createProviderModel(providerId, payload);
toast.success('Model override created successfully!');
} else {
await AdminService.updateProviderModel(
providerId,
adminModelData.id,
payload
);
toast.success('Model updated successfully!');
}
onModelUpdate?.();
onCancel?.();
} catch (error) {
const action = isNewOverride ? 'create' : 'update';
toast.error(`Failed to ${action} model. Please try again.`);
console.error(`Error ${action}ing model:`, error);
} finally {
setIsSubmitting(false);
}
};
const handleClose = () => {
if (!isSubmitting) {
onCancel?.();
}
};
return (
<Dialog open={isOpen} onOpenChange={handleClose}>
<DialogContent className='max-h-[90vh] overflow-y-auto sm:max-w-[600px]'>
<DialogHeader>
<DialogTitle className='flex items-center gap-2'>
<Edit3 className='h-5 w-5' />
{isNewOverride ? 'Create Model Override' : 'Edit Model Override'}
</DialogTitle>
<DialogDescription>
{isNewOverride
? `Create an override for &quot;${model.name}&quot;`
: `Update the model override for &quot;${model.name}&quot;`}
</DialogDescription>
</DialogHeader>
<Form {...form}>
<form onSubmit={form.handleSubmit(onSubmit)} className='space-y-4'>
<div className='grid grid-cols-1 gap-4 sm:grid-cols-2'>
<FormField
control={form.control}
name='name'
render={({ field }) => (
<FormItem>
<FormLabel>Display Name *</FormLabel>
<FormControl>
<Input
placeholder='e.g., GPT-4'
{...field}
className='w-full'
/>
</FormControl>
<FormDescription>
Custom display name for the model
</FormDescription>
<FormMessage />
</FormItem>
)}
/>
<FormField
control={form.control}
name='context_length'
render={({ field }) => (
<FormItem>
<FormLabel>Context Length *</FormLabel>
<FormControl>
<Input
type='number'
min='0'
placeholder='4096'
value={field.value ?? ''}
onChange={(e) => {
const value = e.target.value;
field.onChange(
value === '' ? 0 : parseInt(value, 10) || 0
);
}}
onBlur={(e) => {
const value = parseInt(e.target.value, 10);
field.onChange(
Number.isNaN(value) ? 0 : Math.max(0, value)
);
}}
className='w-full'
/>
</FormControl>
<FormDescription>
Maximum context window size
</FormDescription>
<FormMessage />
</FormItem>
)}
/>
</div>
<FormField
control={form.control}
name='description'
render={({ field }) => (
<FormItem>
<FormLabel>Description</FormLabel>
<FormControl>
<Textarea
placeholder='Brief description of the model...'
{...field}
rows={3}
className='w-full'
/>
</FormControl>
<FormDescription>
Optional description or notes about the model
</FormDescription>
<FormMessage />
</FormItem>
)}
/>
<div className='grid grid-cols-1 gap-4 sm:grid-cols-2'>
<FormField
control={form.control}
name='prompt'
render={({ field }) => (
<FormItem>
<FormLabel>Input Cost (per 1M tokens) *</FormLabel>
<FormControl>
<Input
type='number'
step='0.00001'
min='0'
placeholder='5.00000'
value={field.value ?? ''}
onChange={(e) => {
const value = e.target.value;
field.onChange(value === '' ? 0 : parseFloat(value));
}}
onBlur={(e) => {
const value = parseFloat(e.target.value);
field.onChange(roundToFiveDecimals(value));
}}
className='w-full'
/>
</FormControl>
<FormDescription>
Cost in USD per 1,000,000 input tokens (max 5 decimals)
</FormDescription>
<FormMessage />
</FormItem>
)}
/>
<FormField
control={form.control}
name='completion'
render={({ field }) => (
<FormItem>
<FormLabel>Output Cost (per 1M tokens) *</FormLabel>
<FormControl>
<Input
type='number'
step='0.00001'
min='0'
placeholder='15.00000'
value={field.value ?? ''}
onChange={(e) => {
const value = e.target.value;
field.onChange(value === '' ? 0 : parseFloat(value));
}}
onBlur={(e) => {
const value = parseFloat(e.target.value);
field.onChange(roundToFiveDecimals(value));
}}
className='w-full'
/>
</FormControl>
<FormDescription>
Cost in USD per 1,000,000 output tokens (max 5 decimals)
</FormDescription>
<FormMessage />
</FormItem>
)}
/>
</div>
<FormField
control={form.control}
name='enabled'
render={({ field }) => (
<FormItem className='flex flex-row items-center justify-between rounded-lg border p-4'>
<div className='space-y-0.5'>
<FormLabel className='text-base'>Model Enabled</FormLabel>
<FormDescription>
Enable or disable this model override
</FormDescription>
</div>
<FormControl>
<Switch
checked={field.value}
onCheckedChange={field.onChange}
/>
</FormControl>
</FormItem>
)}
/>
<div className='flex justify-end gap-2 pt-4'>
<Button
type='button'
variant='outline'
onClick={handleClose}
disabled={isSubmitting}
>
Cancel
</Button>
<Button type='submit' disabled={isSubmitting}>
{isSubmitting ? (
<>
<Loader2 className='mr-2 h-4 w-4 animate-spin' />
{isNewOverride ? 'Creating...' : 'Updating...'}
</>
) : isNewOverride ? (
'Create Override'
) : (
'Update Model'
)}
</Button>
</div>
</form>
</Form>
</DialogContent>
</Dialog>
);
}

View File

@@ -197,7 +197,11 @@ export function ModelTester({ models }: ModelTesterProps) {
testModelMutation.mutate(request);
};
const enabledModels = models.filter((model) => model.isEnabled);
const enabledModels = Array.from(
new Map(
models.filter((model) => model.isEnabled).map((m) => [m.id, m])
).values()
);
const credentials = selectedModel ? getModelCredentials(selectedModel) : null;
return (

View File

@@ -124,12 +124,14 @@ export function ModelsPage() {
>
Basic Testing
</TabsTrigger>
{/*
<TabsTrigger
value='test-api'
className='h-9 snap-start px-2 text-[13px] sm:h-10 sm:px-2.5 sm:text-sm'
>
API Endpoints
</TabsTrigger>
*/}
</TabsList>
<TabsContent value='manage' className='mt-0'>

View File

@@ -1,6 +1,6 @@
'use client';
import { useCallback, useState } from 'react';
import { useState } from 'react';
import Image from 'next/image';
import { Copy, Loader2, Zap, KeyRound } from 'lucide-react';
import { toast } from 'sonner';
@@ -10,7 +10,6 @@ import { Input } from '@/components/ui/input';
import { Textarea } from '@/components/ui/textarea';
import { Label } from '@/components/ui/label';
import { Badge } from '@/components/ui/badge';
import { Separator } from '@/components/ui/separator';
import { Tabs, TabsContent, TabsList, TabsTrigger } from '@/components/ui/tabs';
interface RoutstrCreateKeySectionProps {

View File

@@ -32,9 +32,18 @@ interface SettingsData {
onion_url?: string;
cashu_mints?: string[];
relays?: string[];
receive_ln_address?: string;
min_payout_sat?: number;
payout_interval_seconds?: number;
[key: string]: unknown;
}
const PAYOUT_KEYS = [
'receive_ln_address',
'min_payout_sat',
'payout_interval_seconds',
] as const;
const HANDLED_KEYS = [
'name',
'description',
@@ -48,6 +57,7 @@ const HANDLED_KEYS = [
'admin_password',
'id',
'updated_at',
...PAYOUT_KEYS,
];
const IGNORED_KEYS = [
@@ -369,6 +379,7 @@ export function AdminSettings() {
const cashuMintsChanged = hasFieldChanged('cashu_mints');
const relaysChanged = hasFieldChanged('relays');
const analyticsSharingChanged = hasFieldChanged('enable_analytics_sharing');
const payoutChanged = PAYOUT_KEYS.some(hasFieldChanged);
const advancedKeys = Object.keys(settings).filter(
(key) => !HANDLED_KEYS.includes(key) && !IGNORED_KEYS.includes(key)
);
@@ -401,8 +412,44 @@ export function AdminSettings() {
setNewRelay('');
};
const resetAnalyticsSharing = () => resetFields(['enable_analytics_sharing']);
const resetPayout = () => resetFields([...PAYOUT_KEYS]);
const resetAdvanced = () => resetFields(advancedKeys);
const payoutFields: ReadonlyArray<{
key: (typeof PAYOUT_KEYS)[number];
label: string;
placeholder: string;
type: 'text' | 'number';
helpText: string;
min?: number;
}> = [
{
key: 'receive_ln_address',
label: 'Lightning Receive Address',
placeholder: 'you@walletofsatoshi.com or LNURL',
type: 'text',
helpText:
'Lightning address (or LNURL) profits are paid out to. Leave empty to disable periodic payouts.',
},
{
key: 'min_payout_sat',
label: 'Minimum Payout (sat)',
placeholder: '210',
type: 'number',
min: 1,
helpText:
'Wallet payouts only fire when at least this many satoshis are available. Must be > 0.',
},
{
key: 'payout_interval_seconds',
label: 'Payout Interval (seconds)',
placeholder: '900',
type: 'number',
min: 1,
helpText: 'How often the payout loop wakes up to check balances.',
},
];
if (loading) {
return (
<div className='space-y-4'>
@@ -620,6 +667,89 @@ export function AdminSettings() {
) : null}
</Card>
{/* Lightning Payout Settings */}
<Card>
<CardHeader>
<CardTitle>Lightning Payout Settings</CardTitle>
<CardDescription>
Tune how node profit is paid out over Lightning. Amounts must be
positive and above your wallet&apos;s minimum-invoice constraints.
</CardDescription>
</CardHeader>
<CardContent className='space-y-4'>
{payoutFields.map((field) => {
const value = settings[field.key];
if (field.type === 'number') {
return (
<div key={field.key} className='space-y-2'>
<Label htmlFor={field.key}>{field.label}</Label>
<Input
id={field.key}
type='number'
min={field.min}
value={
typeof value === 'number'
? value
: value === undefined || value === null
? ''
: Number(value)
}
placeholder={field.placeholder}
onChange={(e) => {
const raw = e.target.value;
if (raw === '') {
handleInputChange(field.key, undefined);
} else {
const parsed = Number(raw);
handleInputChange(
field.key,
Number.isFinite(parsed) ? parsed : undefined
);
}
}}
/>
<p className='text-muted-foreground text-xs'>
{field.helpText}
</p>
</div>
);
}
return (
<div key={field.key} className='space-y-2'>
<Label htmlFor={field.key}>{field.label}</Label>
<Input
id={field.key}
value={(value as string) || ''}
placeholder={field.placeholder}
onChange={(e) =>
handleInputChange(field.key, e.target.value)
}
/>
<p className='text-muted-foreground text-xs'>
{field.helpText}
</p>
</div>
);
})}
</CardContent>
{payoutChanged ? (
<CardFooter className='justify-start'>
<div className='flex w-full flex-col gap-2 sm:w-auto sm:flex-row sm:items-center'>
<Button
variant='outline'
onClick={resetPayout}
disabled={loading || saving}
>
Cancel
</Button>
<Button onClick={handleSave} disabled={loading || saving}>
{saving ? 'Saving...' : 'Save'}
</Button>
</div>
</CardFooter>
) : null}
</Card>
{/* Relays */}
<Card>
<CardHeader>

View File

@@ -0,0 +1,265 @@
'use client';
import * as React from 'react';
import { useState, useEffect, useCallback } from 'react';
import {
AdminService,
type CliTokenListItem,
type CliTokenCreated,
} from '@/lib/api/services/admin';
import {
Card,
CardContent,
CardHeader,
CardTitle,
CardDescription,
} from '@/components/ui/card';
import { Button } from '@/components/ui/button';
import { Input } from '@/components/ui/input';
import { Label } from '@/components/ui/label';
import { Skeleton } from '@/components/ui/skeleton';
import { Alert, AlertDescription } from '@/components/ui/alert';
import { AlertCircle, Copy, Trash2, Check } from 'lucide-react';
import { toast } from 'sonner';
function formatTs(ts: number | null): string {
if (!ts) return '—';
return new Date(ts * 1000).toLocaleString();
}
export function CliTokensSettings(): React.ReactElement {
const [tokens, setTokens] = useState<CliTokenListItem[]>([]);
const [loading, setLoading] = useState(true);
const [error, setError] = useState<string | null>(null);
const [name, setName] = useState('');
const [expiresInDays, setExpiresInDays] = useState<string>('');
const [creating, setCreating] = useState(false);
const [newToken, setNewToken] = useState<CliTokenCreated | null>(null);
const [copied, setCopied] = useState(false);
const loadTokens = useCallback(async (): Promise<void> => {
setLoading(true);
setError(null);
try {
const data = await AdminService.listCliTokens();
setTokens(data);
} catch (err: unknown) {
const message =
err instanceof Error ? err.message : 'Failed to load tokens';
setError(message);
} finally {
setLoading(false);
}
}, []);
useEffect(() => {
void loadTokens();
}, [loadTokens]);
async function handleCreate(): Promise<void> {
const trimmed = name.trim();
if (!trimmed) {
toast.error('Name is required');
return;
}
const days = expiresInDays.trim()
? Number.parseInt(expiresInDays.trim(), 10)
: undefined;
if (days !== undefined && (Number.isNaN(days) || days <= 0)) {
toast.error('Expiry must be a positive number of days');
return;
}
setCreating(true);
try {
const created = await AdminService.createCliToken(trimmed, days);
setNewToken(created);
setName('');
setExpiresInDays('');
await loadTokens();
toast.success('Token created. Copy it now — it will not be shown again.');
} catch (err: unknown) {
const message =
err instanceof Error ? err.message : 'Failed to create token';
toast.error(message);
} finally {
setCreating(false);
}
}
async function handleRevoke(id: string): Promise<void> {
if (
!confirm('Revoke this token? Any CLI/agent using it will lose access.')
) {
return;
}
try {
await AdminService.revokeCliToken(id);
await loadTokens();
toast.success('Token revoked');
} catch (err: unknown) {
const message =
err instanceof Error ? err.message : 'Failed to revoke token';
toast.error(message);
}
}
async function handleCopy(): Promise<void> {
if (!newToken) return;
await navigator.clipboard.writeText(newToken.token);
setCopied(true);
setTimeout(() => setCopied(false), 2000);
}
return (
<div className='space-y-6'>
<Card>
<CardHeader>
<CardTitle>Create CLI Token</CardTitle>
<CardDescription>
Generate a long-lived bearer token for the Routstr CLI or AI agents.
Use this token in <code>~/.routstr/config.json</code> or with{' '}
<code>routstr init --token &lt;token&gt;</code>.
</CardDescription>
</CardHeader>
<CardContent className='space-y-4'>
{newToken && (
<Alert className='border-green-500/50 bg-green-500/10'>
<AlertDescription className='space-y-3'>
<div className='font-medium text-green-700 dark:text-green-400'>
Token created. Copy it now it will not be shown again.
</div>
<div className='flex items-center gap-2'>
<code className='bg-muted flex-1 rounded px-3 py-2 text-xs break-all'>
{newToken.token}
</code>
<Button
type='button'
variant='outline'
size='sm'
onClick={handleCopy}
>
{copied ? (
<Check className='h-4 w-4' />
) : (
<Copy className='h-4 w-4' />
)}
</Button>
</div>
<Button
type='button'
variant='ghost'
size='sm'
onClick={() => setNewToken(null)}
>
Dismiss
</Button>
</AlertDescription>
</Alert>
)}
<div className='grid grid-cols-1 gap-4 md:grid-cols-2'>
<div className='space-y-2'>
<Label htmlFor='cli-token-name'>Name</Label>
<Input
id='cli-token-name'
placeholder='e.g. dev-laptop, ci-runner'
value={name}
onChange={(e) => setName(e.target.value)}
disabled={creating}
/>
</div>
<div className='space-y-2'>
<Label htmlFor='cli-token-expiry'>
Expires in days (optional)
</Label>
<Input
id='cli-token-expiry'
type='number'
min='1'
placeholder='Never expires if blank'
value={expiresInDays}
onChange={(e) => setExpiresInDays(e.target.value)}
disabled={creating}
/>
</div>
</div>
<Button onClick={handleCreate} disabled={creating || !name.trim()}>
{creating ? 'Creating…' : 'Create Token'}
</Button>
</CardContent>
</Card>
<Card>
<CardHeader>
<CardTitle>Active Tokens</CardTitle>
<CardDescription>
Tokens authorize CLI/agent calls to admin endpoints. Revoke any
token that may have been exposed.
</CardDescription>
</CardHeader>
<CardContent>
{error && (
<Alert variant='destructive' className='mb-4'>
<AlertCircle className='h-4 w-4' />
<AlertDescription>{error}</AlertDescription>
</Alert>
)}
{loading ? (
<div className='space-y-2'>
<Skeleton className='h-12 w-full' />
<Skeleton className='h-12 w-full' />
</div>
) : tokens.length === 0 ? (
<p className='text-muted-foreground text-sm'>
No tokens yet. Create one above.
</p>
) : (
<div className='overflow-x-auto'>
<table className='w-full text-sm'>
<thead>
<tr className='text-muted-foreground border-b text-left'>
<th className='py-2 pr-4 font-medium'>Name</th>
<th className='py-2 pr-4 font-medium'>Token</th>
<th className='py-2 pr-4 font-medium'>Created</th>
<th className='py-2 pr-4 font-medium'>Last used</th>
<th className='py-2 pr-4 font-medium'>Expires</th>
<th className='py-2 font-medium'></th>
</tr>
</thead>
<tbody>
{tokens.map((t) => (
<tr key={t.id} className='border-b last:border-0'>
<td className='py-2 pr-4'>{t.name}</td>
<td className='py-2 pr-4 font-mono text-xs'>
{t.token_preview}
</td>
<td className='text-muted-foreground py-2 pr-4'>
{formatTs(t.created_at)}
</td>
<td className='text-muted-foreground py-2 pr-4'>
{formatTs(t.last_used_at)}
</td>
<td className='text-muted-foreground py-2 pr-4'>
{t.expires_at ? formatTs(t.expires_at) : 'Never'}
</td>
<td className='py-2'>
<Button
type='button'
variant='ghost'
size='sm'
onClick={() => void handleRevoke(t.id)}
>
<Trash2 className='h-4 w-4' />
</Button>
</td>
</tr>
))}
</tbody>
</table>
</div>
)}
</CardContent>
</Card>
</div>
);
}

View File

@@ -0,0 +1,282 @@
'use client';
import { useQuery, useQueryClient } from '@tanstack/react-query';
import {
AlertTriangleIcon,
CheckCircle2Icon,
ExternalLinkIcon,
InfoIcon,
Loader2Icon,
RefreshCwIcon,
} from 'lucide-react';
import { ConfigurationService } from '@/lib/api/services/configuration';
import { Button } from '@/components/ui/button';
import {
Popover,
PopoverContent,
PopoverTrigger,
} from '@/components/ui/popover';
import { cn } from '@/lib/utils';
import {
deriveStatus,
formatReleaseDate,
formatVersionLabel,
parseVersion,
type StatusKind,
} from '@/lib/utils/version';
interface NodeInfo {
version?: string;
}
interface GithubRelease {
tag_name: string;
name?: string;
html_url: string;
published_at?: string;
body?: string;
}
const NODE_QUERY_KEY = ['node-version'] as const;
const RELEASE_QUERY_KEY = ['routstr-latest-release'] as const;
const THIRTY_MINUTES = 30 * 60 * 1000;
const GITHUB_RELEASES_URL = `https://api.github.com/repos/Routstr/routstr-core/releases/latest`;
const RELEASES_PAGE_URL = `https://github.com/Routstr/routstr-core/releases`;
async function fetchNodeInfo(): Promise<NodeInfo> {
const baseUrl = ConfigurationService.getLocalBaseUrl().replace(/\/+$/, '');
const response = await fetch(`${baseUrl}/v1/info`, {
headers: { 'Content-Type': 'application/json' },
});
if (!response.ok) {
throw new Error('Unable to load node info');
}
return (await response.json()) as NodeInfo;
}
async function fetchLatestRelease(): Promise<GithubRelease | null> {
const response = await fetch(GITHUB_RELEASES_URL, {
headers: { Accept: 'application/vnd.github+json' },
});
if (response.status === 403 || response.status === 404) {
return null;
}
if (!response.ok) {
throw new Error(`GitHub responded ${response.status}`);
}
return (await response.json()) as GithubRelease;
}
function pickColorClass(status: StatusKind): string {
if (status === 'outdated') return 'text-amber-600 dark:text-amber-400';
if (status === 'unknown') return 'text-muted-foreground';
if (status === 'ahead' || status === 'commit-drift') {
return 'text-sky-600 dark:text-sky-400';
}
return 'text-emerald-600 dark:text-emerald-400';
}
function renderStatusIcon(status: StatusKind, className: string) {
if (status === 'outdated') {
return <AlertTriangleIcon className={className} />;
}
if (status === 'commit-drift' || status === 'ahead' || status === 'unknown') {
return <InfoIcon className={className} />;
}
return <CheckCircle2Icon className={className} />;
}
function describeStatus(status: StatusKind): string {
if (status === 'outdated') return 'A newer release is available.';
if (status === 'commit-drift') {
return 'Running release version on a non-release commit.';
}
if (status === 'ahead') {
return 'Running ahead of the latest published release.';
}
if (status === 'current') return 'Up to date with the latest release.';
return 'Version status unavailable.';
}
interface VersionStatusProps {
variant?: 'expanded' | 'compact';
className?: string;
}
export function VersionStatus({
variant = 'expanded',
className,
}: VersionStatusProps) {
const queryClient = useQueryClient();
const nodeQuery = useQuery({
queryKey: NODE_QUERY_KEY,
queryFn: fetchNodeInfo,
staleTime: THIRTY_MINUTES,
retry: 1,
});
const releaseQuery = useQuery({
queryKey: RELEASE_QUERY_KEY,
queryFn: fetchLatestRelease,
staleTime: THIRTY_MINUTES,
refetchInterval: THIRTY_MINUTES,
refetchOnWindowFocus: false,
retry: 1,
});
const currentVersion = parseVersion(nodeQuery.data?.version);
const latestVersion = parseVersion(releaseQuery.data?.tag_name);
const status = deriveStatus(currentVersion, latestVersion);
const isRefreshing = releaseQuery.isFetching || nodeQuery.isFetching;
const handleRefresh = async (): Promise<void> => {
await Promise.all([
queryClient.invalidateQueries({ queryKey: NODE_QUERY_KEY }),
queryClient.invalidateQueries({ queryKey: RELEASE_QUERY_KEY }),
]);
};
const colorClass = pickColorClass(status);
const versionLabel = currentVersion
? formatVersionLabel(currentVersion)
: nodeQuery.isLoading
? '…'
: 'unknown';
if (!nodeQuery.data && nodeQuery.isLoading && variant === 'expanded') {
return null;
}
const statusDescription = describeStatus(status);
const ariaLabel = `Node version ${versionLabel}. ${statusDescription} Click for details.`;
const releaseRateLimited = releaseQuery.data === null;
return (
<Popover>
<PopoverTrigger asChild>
<button
type='button'
onClick={(e) => e.stopPropagation()}
className={cn(
'hover:bg-accent/40 inline-flex items-center gap-1 rounded-md px-1 py-0.5 font-mono text-[10px] leading-tight transition-colors',
colorClass,
className
)}
title='View version details'
aria-label={ariaLabel}
>
{renderStatusIcon(status, 'h-3 w-3 shrink-0')}
<span className='truncate'>{versionLabel}</span>
</button>
</PopoverTrigger>
<PopoverContent
side='bottom'
align='start'
sideOffset={6}
className='w-72 p-3'
>
<div className='flex items-start justify-between gap-2'>
<div className='min-w-0 space-y-0.5'>
<div className='flex items-center gap-1.5 text-sm font-medium'>
{renderStatusIcon(status, cn('h-4 w-4 shrink-0', colorClass))}
Node Version
</div>
<p className='text-muted-foreground text-xs leading-snug'>
{statusDescription}
</p>
</div>
<Button
type='button'
variant='outline'
size='icon'
className='h-7 w-7 shrink-0'
onClick={handleRefresh}
disabled={isRefreshing}
title='Check for latest release'
>
{isRefreshing ? (
<Loader2Icon className='h-3.5 w-3.5 animate-spin' />
) : (
<RefreshCwIcon className='h-3.5 w-3.5' />
)}
<span className='sr-only'>Check for latest release</span>
</Button>
</div>
<div className='mt-3 space-y-2'>
<div className='border-border/60 bg-card/30 grid gap-1.5 rounded-md border p-2'>
<div className='flex items-center justify-between gap-3'>
<span className='text-muted-foreground text-[10px] tracking-wide uppercase'>
Current
</span>
<span className='font-mono text-xs'>{versionLabel}</span>
</div>
{currentVersion?.commit ? (
<div className='flex items-center justify-between gap-3'>
<span className='text-muted-foreground text-[10px] tracking-wide uppercase'>
Commit
</span>
<code className='font-mono text-[11px]'>
{currentVersion.commit}
</code>
</div>
) : null}
</div>
<div className='border-border/60 bg-card/30 grid gap-1.5 rounded-md border p-2'>
<div className='flex items-center justify-between gap-3'>
<span className='text-muted-foreground text-[10px] tracking-wide uppercase'>
Latest release
</span>
<span className='font-mono text-xs'>
{releaseQuery.isLoading
? 'loading…'
: releaseQuery.isError
? 'unavailable'
: releaseRateLimited
? 'rate-limited'
: (releaseQuery.data?.tag_name ?? 'unknown')}
</span>
</div>
{releaseQuery.data?.published_at ? (
<div className='flex items-center justify-between gap-3'>
<span className='text-muted-foreground text-[10px] tracking-wide uppercase'>
Published
</span>
<span className='text-[11px]'>
{formatReleaseDate(releaseQuery.data.published_at)}
</span>
</div>
) : null}
</div>
{releaseQuery.isError ? (
<p className='text-muted-foreground text-[11px]'>
Failed to fetch latest release from GitHub.
</p>
) : releaseRateLimited ? (
<p className='text-muted-foreground text-[11px]'>
GitHub rate limit reached. Try again later.
</p>
) : null}
</div>
<div className='border-border/60 mt-3 border-t pt-2'>
<a
href={releaseQuery.data?.html_url ?? RELEASES_PAGE_URL}
target='_blank'
rel='noopener noreferrer'
className='text-primary inline-flex items-center gap-1 text-xs hover:underline'
>
View release changelog
<ExternalLinkIcon className='h-3 w-3' />
</a>
</div>
</PopoverContent>
</Popover>
);
}

View File

@@ -76,6 +76,7 @@ export const AdminModelSchema = z.object({
canonical_slug: z.string().nullable().optional(),
alias_ids: z.array(z.string()).nullable().optional(),
enabled: z.boolean().default(true),
forwarded_model_id: z.string().nullable().optional(),
});
export const ProviderModelsSchema = z.object({
@@ -890,19 +891,42 @@ export class AdminService {
type?: string,
status?: string,
search?: string,
limit: number = 100
source?: string,
limit: number = 50,
offset: number = 0
): Promise<TransactionsResponse> {
const params = new URLSearchParams();
if (type) params.append('type', type);
if (status) params.append('status', status);
if (search) params.append('search', search);
if (source) params.append('source', source);
params.append('limit', limit.toString());
params.append('offset', offset.toString());
return await apiClient.get<TransactionsResponse>(
`/admin/api/transactions?${params.toString()}`
);
}
static async getLightningInvoices(
status?: string,
purpose?: string,
search?: string,
limit: number = 50,
offset: number = 0
): Promise<LightningInvoicesResponse> {
const params = new URLSearchParams();
if (status) params.append('status', status);
if (purpose) params.append('purpose', purpose);
if (search) params.append('search', search);
params.append('limit', limit.toString());
params.append('offset', offset.toString());
return await apiClient.get<LightningInvoicesResponse>(
`/admin/api/lightning-invoices?${params.toString()}`
);
}
static async createProviderAccountByType(providerType: string): Promise<{
ok: boolean;
account_data: Record<string, unknown>;
@@ -961,6 +985,45 @@ export class AdminService {
balance_data: number | null | Record<string, unknown>;
}>(`/admin/api/upstream-providers/${providerId}/balance`);
}
// ── CLI Tokens ──
static async listCliTokens(): Promise<CliTokenListItem[]> {
return await apiClient.get<CliTokenListItem[]>('/admin/api/cli-tokens');
}
static async createCliToken(
name: string,
expiresInDays?: number
): Promise<CliTokenCreated> {
return await apiClient.post<CliTokenCreated>('/admin/api/cli-tokens', {
name,
expires_in_days: expiresInDays ?? null,
});
}
static async revokeCliToken(tokenId: string): Promise<{ ok: boolean }> {
return await apiClient.delete<{ ok: boolean }>(
`/admin/api/cli-tokens/${encodeURIComponent(tokenId)}`
);
}
}
export interface CliTokenListItem {
id: string;
name: string;
token_preview: string;
created_at: number;
last_used_at: number | null;
expires_at: number | null;
}
export interface CliTokenCreated {
id: string;
name: string;
token: string;
created_at: number;
expires_at: number | null;
}
export const TemporaryBalanceSchema = z.object({
@@ -1134,9 +1197,30 @@ export interface Transaction {
created_at: number;
collected: boolean;
swept: boolean;
source: 'x-cashu' | 'apikey';
api_key_hashed_key?: string;
}
export interface TransactionsResponse {
transactions: Transaction[];
total: number;
}
export interface LightningInvoice {
id: string;
bolt11: string;
amount_sats: number;
description: string;
payment_hash: string;
status: 'pending' | 'paid' | 'expired' | 'cancelled';
api_key_hash: string | null;
purpose: 'create' | 'topup';
created_at: number;
expires_at: number;
paid_at: number | null;
}
export interface LightningInvoicesResponse {
invoices: LightningInvoice[];
total: number;
}

72
ui/lib/utils/version.ts Normal file
View File

@@ -0,0 +1,72 @@
export interface ParsedVersion {
raw: string;
base: string;
parts: readonly number[];
commit: string | null;
}
export type StatusKind =
| 'unknown'
| 'outdated'
| 'ahead'
| 'commit-drift'
| 'current';
export function parseVersion(
raw: string | undefined | null
): ParsedVersion | null {
if (!raw) return null;
const trimmed = raw.trim().replace(/^v/i, '');
if (!trimmed) return null;
const [base, commitPart] = trimmed.split('+', 2);
const parts = (base ?? '')
.split('.')
.map((segment) => Number.parseInt(segment, 10))
.filter((value) => Number.isFinite(value));
if (parts.length === 0) return null;
return {
raw,
base: base ?? '',
parts,
commit: commitPart ?? null,
};
}
export function compareVersionParts(
a: readonly number[],
b: readonly number[]
): number {
const length = Math.max(a.length, b.length);
for (let i = 0; i < length; i += 1) {
const diff = (a[i] ?? 0) - (b[i] ?? 0);
if (diff !== 0) return diff;
}
return 0;
}
export function deriveStatus(
current: ParsedVersion | null,
latest: ParsedVersion | null
): StatusKind {
if (!current || !latest) return 'unknown';
const cmp = compareVersionParts(current.parts, latest.parts);
if (cmp < 0) return 'outdated';
if (cmp > 0) return 'ahead';
return current.commit ? 'commit-drift' : 'current';
}
export function formatReleaseDate(iso: string | undefined): string | null {
if (!iso) return null;
const date = new Date(iso);
if (Number.isNaN(date.getTime())) return null;
return date.toLocaleDateString(undefined, {
year: 'numeric',
month: 'short',
day: 'numeric',
});
}
export function formatVersionLabel(version: ParsedVersion | null): string {
if (!version) return 'unknown';
return version.raw.startsWith('v') ? version.raw : `v${version.raw}`;
}

View File

@@ -6,6 +6,9 @@ const nextConfig: NextConfig = {
images: {
unoptimized: true,
},
turbopack: {
root: __dirname,
},
};
export default nextConfig;

View File

@@ -1,5 +1,6 @@
{
"name": "routstr-service",
"packageManager": "pnpm@10.15.0",
"version": "0.1.0",
"private": true,
"scripts": {
@@ -88,5 +89,11 @@
"prettier-plugin-tailwindcss": "^0.7.2",
"tailwindcss": "^4.2.0",
"typescript": "^5.9.3"
},
"pnpm": {
"onlyBuiltDependencies": [
"sharp",
"unrs-resolver"
]
}
}

3607
uv.lock generated

File diff suppressed because it is too large Load Diff