Compare commits

..
56 Commits
Author SHA1 Message Date
LucasDLTG bdbd9413fb fix: fix and check all llm endpoints
CI / Rust CI (push) Failing after 1m54s
2026-07-14 23:11:20 +02:00
LucasDLTG f1b8c310c4 feat: complete docs 2026-07-14 13:18:37 +02:00
LucasDLTG 387e0a0cfb feat: add unload llm route 2026-07-14 12:16:42 +02:00
LucasDLTG 67312cb7a4 fix: put auth errors in good place 2026-07-14 11:20:06 +02:00
LucasDLTG 63e57153dd feat: prepare tts 2026-07-14 10:09:47 +02:00
LucasDLTG 77bd729d95 fix: streaming
CI / Rust CI (push) Successful in 5m19s
2026-06-05 18:16:48 +02:00
LucasDLTG 28352d8bdb reafctor: all code without stream 2026-06-02 18:30:39 +02:00
LucasDLTG 974be437af feat: use json in http error 2026-05-21 22:31:03 +02:00
LucasDLTG 739539309a refactor: postgres files 2026-05-21 22:27:03 +02:00
LucasDLTG 0d31ade1f6 feat: conversation retrieveing 2026-05-21 19:58:42 +02:00
LucasDLTG f01bc74059 feat: log user message 2026-05-13 12:39:02 +02:00
LucasDLTG c976db7063 feat: return conversation id 2026-05-11 19:04:14 +02:00
LucasDLTG b5ef7e4881 feat: generate title 2026-05-11 18:28:04 +02:00
LucasDLTG 1bf4cf2ca3 feat: create conversation 2026-05-11 15:12:15 +02:00
LucasDLTG 5aa7a6e7e1 feat: db errors 2026-05-11 11:33:12 +02:00
LucasDLTG b77715268d refactor: ollama errors centralized 2026-05-11 11:16:16 +02:00
LucasDLTG 7bc489939f feat: prepare log chat 2026-05-11 10:54:29 +02:00
LucasDLTG b7e5e04fc0 feat: get roles associted to api key 2026-05-08 13:11:59 +02:00
LucasDLTG eaf197f86e feat: secure genrate key endpoint
CI / Rust CI (push) Successful in 2m9s
Publish & Deploy / Build and Push to Registry (push) Successful in 28s
Publish & Deploy / Deploy via SSH (push) Successful in 18s
2026-05-07 20:59:01 +02:00
LucasDLTG 7b98fab2c8 fix: add cors option http
CI / Rust CI (push) Successful in 2m8s
Publish & Deploy / Build and Push to Registry (push) Successful in 28s
Publish & Deploy / Deploy via SSH (push) Successful in 17s
2026-05-07 20:12:24 +02:00
LucasDLTG d7ab571bbe fix: release ci with sqlx
CI / Rust CI (push) Successful in 2m8s
Publish & Deploy / Build and Push to Registry (push) Successful in 27s
Publish & Deploy / Deploy via SSH (push) Successful in 7s
2026-05-07 19:56:17 +02:00
LucasDLTG 87fcd522c4 fix: release ci with sqlx
CI / Rust CI (push) Successful in 2m9s
Publish & Deploy / Build and Push to Registry (push) Failing after 14s
Publish & Deploy / Deploy via SSH (push) Has been skipped
2026-05-07 19:50:37 +02:00
LucasDLTG 5884d840e7 fix: cargo audit no fatal
CI / Rust CI (push) Successful in 5m12s
Publish & Deploy / Build and Push to Registry (push) Failing after 14s
Publish & Deploy / Deploy via SSH (push) Has been skipped
2026-05-07 19:43:20 +02:00
LucasDLTG a631b8cd1e feat: add .sqlx 2026-05-07 19:39:50 +02:00
LucasDLTG 8bc5200f77 feat: readme 2026-05-07 19:36:44 +02:00
LucasDLTG 8878dbb454 fix: pre commit
CI / Rust CI (push) Failing after 1m12s
2026-05-07 19:34:43 +02:00
LucasDLTG ecf275dedd feat: readme 2026-05-07 19:33:25 +02:00
LucasDLTG a6b9ee5b09 fix: ci with sqlx
CI / Rust CI (push) Failing after 1m10s
Publish & Deploy / Build and Push to Registry (push) Failing after 13s
Publish & Deploy / Deploy via SSH (push) Has been skipped
2026-05-07 19:17:01 +02:00
LucasDLTG 2e81132f2d fix: use nex table name
CI / Rust CI (push) Failing after 1m12s
Publish & Deploy / Build and Push to Registry (push) Failing after 1m48s
Publish & Deploy / Deploy via SSH (push) Has been skipped
2026-05-07 19:07:14 +02:00
LucasDLTG c534299c8e feat: update last time used api key 2026-05-07 16:15:27 +02:00
LucasDLTG 4b428ec32a feat: add api key verification 2026-05-07 14:13:29 +02:00
LucasDLTG 752373c7b7 feat: add api key generation endpint + role authorization 2026-05-07 12:32:29 +02:00
LucasDLTG d7ddc087a6 feat: add user in db 2026-05-06 15:22:49 +02:00
LucasDLTG 7f77bfef4e fix: port
CI / Rust CI (push) Successful in 1m37s
Publish & Deploy / Build and Push to Registry (push) Successful in 13s
Publish & Deploy / Deploy via SSH (push) Successful in 17s
2026-04-28 18:49:28 +02:00
LucasDLTG 4e64aab0a5 fix: format + tracing
CI / Rust CI (push) Successful in 1m39s
Publish & Deploy / Build and Push to Registry (push) Successful in 1m37s
Publish & Deploy / Deploy via SSH (push) Successful in 18s
2026-04-28 18:32:40 +02:00
LucasDLTG 4596dfe28b feat: add tracing
CI / Rust CI (push) Successful in 4m41s
2026-04-25 19:55:58 +02:00
LucasDLTG 57a3d04625 feat: add cors
CI / Rust CI (push) Successful in 4m37s
2026-04-20 21:03:18 +02:00
LucasDLTG 81129d6c9c feat: add metadata for api doc
CI / Rust CI (push) Successful in 1m33s
Publish & Deploy / Build and Push to Registry (push) Successful in 22s
Publish & Deploy / Deploy via SSH (push) Successful in 18s
2026-04-11 20:40:38 +02:00
LucasDLTG f4da65dbc2 refactor: middlewares folder
CI / Rust CI (push) Successful in 1m37s
Publish & Deploy / Build and Push to Registry (push) Successful in 24s
Publish & Deploy / Deploy via SSH (push) Successful in 17s
2026-04-11 20:35:36 +02:00
LucasDLTG e1f46c07b4 feat: add doc endpoint 2026-04-11 20:23:11 +02:00
LucasDLTG fc391b5d0e feat: update test
CI / Rust CI (push) Successful in 4m35s
Publish & Deploy / Build and Push to Registry (push) Successful in 1m30s
Publish & Deploy / Deploy via SSH (push) Successful in 18s
2026-04-10 22:26:45 +02:00
LucasDLTG e31cdf131f feat: chat typing 2026-04-10 21:33:52 +02:00
LucasDLTG d5856557b4 feat: strong typing for complete endpoints 2026-04-10 18:27:02 +02:00
LucasDLTG 562d154480 feat: add proper typing for api 2026-04-10 17:41:25 +02:00
LucasDLTG a22560c337 feat: add auto doc via utiopia
CI / Rust CI (push) Successful in 4m29s
2026-04-10 15:45:41 +02:00
LucasDLTG b3ad5249c0 update: .env.example
CI / Rust CI (push) Successful in 1m25s
2026-04-10 14:38:06 +02:00
LucasDLTG 1a99490e22 feat: add streaming
CI / Rust CI (push) Successful in 4m29s
2026-04-10 14:15:27 +02:00
LucasDLTG 686f9ff747 feat: add unload endpoint
CI / Rust CI (push) Successful in 1m24s
2026-04-10 12:48:14 +02:00
LucasDLTG 7e430bcf00 fix auth debug infos only in debug mode
CI / Rust CI (push) Successful in 1m26s
2026-04-10 12:35:00 +02:00
LucasDLTG 1a06577fa3 feat: add load endpoint 2026-04-10 12:29:34 +02:00
LucasDLTG db14d816cb format: separate routes in file system 2026-04-10 12:08:03 +02:00
LucasDLTG 6dc7231309 feat: add chat completions route + test
CI / Rust CI (push) Successful in 4m23s
Publish & Deploy / Build and Push to Registry (push) Successful in 1m26s
Publish & Deploy / Deploy via SSH (push) Successful in 18s
2026-04-10 11:47:56 +02:00
LucasDLTG 70fe9bc4da feat: add clean errors + completion endpoint
CI / Rust CI (push) Successful in 4m12s
2026-04-10 11:18:20 +02:00
LucasDLTG fa09e75a09 feat: use env var for ollama url
CI / Rust CI (push) Successful in 1m11s
Publish & Deploy / Build and Push to Registry (push) Successful in 17s
Publish & Deploy / Deploy via SSH (push) Successful in 18s
2026-04-09 21:17:29 +02:00
LucasDLTG 46e455fa07 feat: put api and ollama in same docker net
CI / Rust CI (push) Successful in 1m13s
Publish & Deploy / Build and Push to Registry (push) Successful in 18s
Publish & Deploy / Deploy via SSH (push) Successful in 17s
2026-04-09 21:01:27 +02:00
LucasDLTG 8d26907c89 feat: add models list endpoint 2026-04-09 20:51:20 +02:00
87 changed files with 7056 additions and 214 deletions
+2
View File
@@ -1,2 +1,4 @@
JWKS_URL=https://auth.iceberg.black/realms/iceberg/protocol/openid-connect/certs JWKS_URL=https://auth.iceberg.black/realms/iceberg/protocol/openid-connect/certs
ISSUER=https://auth.iceberg.black/realms/iceberg ISSUER=https://auth.iceberg.black/realms/iceberg
OLLAMA_URL=...
CORS_ORIGIN=
+4 -2
View File
@@ -32,8 +32,10 @@ jobs:
restore-keys: | restore-keys: |
${{ runner.os }}-cargo- ${{ runner.os }}-cargo-
# - name: Run tests - name: Run tests
# run: cargo test --workspace --verbose env:
SQLX_OFFLINE: true
run: cargo test --workspace --verbose --all
- name: Run Clippy (linter) - name: Run Clippy (linter)
run: cargo clippy --all-targets --all-features -- -D warnings run: cargo clippy --all-targets --all-features -- -D warnings
+3 -2
View File
@@ -3,7 +3,7 @@ name: Publish & Deploy
on: on:
push: push:
tags: tags:
- 'v*' - "v*"
permissions: permissions:
contents: write contents: write
@@ -34,6 +34,8 @@ jobs:
- name: Build and Push - name: Build and Push
uses: https://github.com/docker/build-push-action@v5 uses: https://github.com/docker/build-push-action@v5
env:
SQLX_OFFLINE: true
with: with:
context: . context: .
push: true push: true
@@ -86,4 +88,3 @@ jobs:
REPO_LOWER=$REPO_LOWER TAG=$TAG docker compose up -d REPO_LOWER=$REPO_LOWER TAG=$TAG docker compose up -d
docker image prune -af docker image prune -af
+11 -4
View File
@@ -12,30 +12,37 @@ repos:
- id: clippy - id: clippy
name: clippy name: clippy
entry: cargo clippy -- -D warnings entry: bash -c 'SQLX_OFFLINE=true cargo clippy -- -D warnings'
language: system language: system
types: [rust] types: [rust]
pass_filenames: false pass_filenames: false
stages: [pre-commit] stages: [pre-commit]
# ─────────────── PRE-PUSH (heavy checks) ─────────────── # ─────────────── PRE-PUSH (heavy checks) ───────────────
- id: sqlx-prepare-check
name: sqlx prepare check
entry: cargo sqlx prepare --check
language: system
pass_filenames: false
stages: [pre-push]
- id: test - id: test
name: cargo test name: cargo test
entry: cargo test --all entry: bash -c 'SQLX_OFFLINE=true cargo test --all'
language: system language: system
pass_filenames: false pass_filenames: false
stages: [pre-push] stages: [pre-push]
- id: build - id: build
name: cargo build release name: cargo build release
entry: cargo build --release entry: bash -c 'SQLX_OFFLINE=true cargo build --release'
language: system language: system
pass_filenames: false pass_filenames: false
stages: [pre-push] stages: [pre-push]
- id: audit - id: audit
name: cargo audit name: cargo audit
entry: cargo audit entry: bash -c 'cargo audit || echo "cargo audit failed (non-blocking)"'
language: system language: system
pass_filenames: false pass_filenames: false
stages: [pre-push] stages: [pre-push]
@@ -0,0 +1,15 @@
{
"db_name": "PostgreSQL",
"query": "UPDATE chat.message SET tokens = $1 WHERE id = $2",
"describe": {
"columns": [],
"parameters": {
"Left": [
"Int4",
"Uuid"
]
},
"nullable": []
},
"hash": "426a96fdb055d248f1f73d52dd10014340c6174b1491c1da36a5ba63fa554a9f"
}
@@ -0,0 +1,51 @@
{
"db_name": "PostgreSQL",
"query": "\n SELECT\n u.id AS user_id,\n ak.id AS api_key_id,\n ak.scopes AS \"roles!: Vec<Role>\"\n FROM auth.api_key ak\n JOIN auth.app_user u ON u.id = ak.created_by\n WHERE ak.key_hash = $1\n AND ak.revoked_at IS NULL\n ",
"describe": {
"columns": [
{
"ordinal": 0,
"name": "user_id",
"type_info": "Uuid"
},
{
"ordinal": 1,
"name": "api_key_id",
"type_info": "Uuid"
},
{
"ordinal": 2,
"name": "roles!: Vec<Role>",
"type_info": {
"Custom": {
"name": "auth.role[]",
"kind": {
"Array": {
"Custom": {
"name": "auth.role",
"kind": {
"Enum": [
"user",
"admin"
]
}
}
}
}
}
}
}
],
"parameters": {
"Left": [
"Text"
]
},
"nullable": [
false,
false,
false
]
},
"hash": "52a8daad9aa79a9c391bfa52b92b268f3910e4e3266f029f17f460c9f5362dd4"
}
@@ -0,0 +1,42 @@
{
"db_name": "PostgreSQL",
"query": "\n SELECT id, title, created_at, updated_at\n FROM chat.conversation\n WHERE user_id = $1\n AND ($2::timestamptz IS NULL OR updated_at < $2)\n ORDER BY updated_at DESC\n LIMIT $3\n ",
"describe": {
"columns": [
{
"ordinal": 0,
"name": "id",
"type_info": "Uuid"
},
{
"ordinal": 1,
"name": "title",
"type_info": "Text"
},
{
"ordinal": 2,
"name": "created_at",
"type_info": "Timestamptz"
},
{
"ordinal": 3,
"name": "updated_at",
"type_info": "Timestamptz"
}
],
"parameters": {
"Left": [
"Uuid",
"Timestamptz",
"Int8"
]
},
"nullable": [
false,
true,
false,
false
]
},
"hash": "6ea72b05912c3075a17f36dc71691fa5dde6934c29e96e9d98f4d65db14c1d88"
}
@@ -0,0 +1,22 @@
{
"db_name": "PostgreSQL",
"query": "INSERT INTO chat.conversation (user_id) VALUES ($1) RETURNING id",
"describe": {
"columns": [
{
"ordinal": 0,
"name": "id",
"type_info": "Uuid"
}
],
"parameters": {
"Left": [
"Uuid"
]
},
"nullable": [
false
]
},
"hash": "93e7661faaa9c5a28468b6b5a4bbb0e40510c081e2d4d5e99bdbf5c153112da2"
}
@@ -0,0 +1,14 @@
{
"db_name": "PostgreSQL",
"query": "\n UPDATE auth.api_key\n SET last_used_at = now()\n WHERE id = $1\n ",
"describe": {
"columns": [],
"parameters": {
"Left": [
"Uuid"
]
},
"nullable": []
},
"hash": "ab03067777ea2a8f12c3dec5cf99caf4ad700008d892a33c25408c34bebab83c"
}
@@ -0,0 +1,34 @@
{
"db_name": "PostgreSQL",
"query": "\n INSERT INTO auth.api_key (key_hash, name, created_by, scopes)\n VALUES ($1, $2, $3, $4::auth.role[])\n ",
"describe": {
"columns": [],
"parameters": {
"Left": [
"Text",
"Text",
"Uuid",
{
"Custom": {
"name": "auth.role[]",
"kind": {
"Array": {
"Custom": {
"name": "auth.role",
"kind": {
"Enum": [
"user",
"admin"
]
}
}
}
}
}
}
]
},
"nullable": []
},
"hash": "ac5047be40cf3c686d3a3f877fc7af232e6a9e9010c947422f0497ff76b06c58"
}
@@ -0,0 +1,65 @@
{
"db_name": "PostgreSQL",
"query": "\n SELECT id, parent_id, role as \"role: MessageRole\", content, created_at, tokens\n FROM chat.message\n WHERE conversation_id = $1\n AND ($2::timestamptz IS NULL OR created_at < $2)\n ORDER BY created_at DESC\n LIMIT $3\n ",
"describe": {
"columns": [
{
"ordinal": 0,
"name": "id",
"type_info": "Uuid"
},
{
"ordinal": 1,
"name": "parent_id",
"type_info": "Uuid"
},
{
"ordinal": 2,
"name": "role: MessageRole",
"type_info": {
"Custom": {
"name": "chat.role",
"kind": {
"Enum": [
"user",
"assistant",
"system"
]
}
}
}
},
{
"ordinal": 3,
"name": "content",
"type_info": "Text"
},
{
"ordinal": 4,
"name": "created_at",
"type_info": "Timestamptz"
},
{
"ordinal": 5,
"name": "tokens",
"type_info": "Int4"
}
],
"parameters": {
"Left": [
"Uuid",
"Timestamptz",
"Int8"
]
},
"nullable": [
false,
true,
false,
false,
false,
true
]
},
"hash": "ac5b689f154241e193cff087a552af61028fd0ae9f34c13504950be6742f107f"
}
@@ -0,0 +1,37 @@
{
"db_name": "PostgreSQL",
"query": "\n INSERT INTO chat.message (conversation_id, parent_id, role, content, tokens)\n VALUES ($1, $2, $3, $4, $5)\n RETURNING id\n ",
"describe": {
"columns": [
{
"ordinal": 0,
"name": "id",
"type_info": "Uuid"
}
],
"parameters": {
"Left": [
"Uuid",
"Uuid",
{
"Custom": {
"name": "chat.role",
"kind": {
"Enum": [
"user",
"assistant",
"system"
]
}
}
},
"Text",
"Int4"
]
},
"nullable": [
false
]
},
"hash": "c98b93484092fc57098f4c6eac92af653bcf39aa1bb4d8add80ed069421ee859"
}
@@ -0,0 +1,22 @@
{
"db_name": "PostgreSQL",
"query": "SELECT id FROM chat.conversation WHERE id = $1",
"describe": {
"columns": [
{
"ordinal": 0,
"name": "id",
"type_info": "Uuid"
}
],
"parameters": {
"Left": [
"Uuid"
]
},
"nullable": [
false
]
},
"hash": "db88cb8d5b447e51fa1bd5ef4feed427e828be9029f4496ec63e704f44982f9e"
}
@@ -0,0 +1,15 @@
{
"db_name": "PostgreSQL",
"query": "UPDATE chat.conversation SET title = $1 WHERE id = $2",
"describe": {
"columns": [],
"parameters": {
"Left": [
"Text",
"Uuid"
]
},
"nullable": []
},
"hash": "ea0755640e64de9eb034de9c1d94bf0f09143b2ddf200757b5e8ac53e1e8a993"
}
@@ -0,0 +1,14 @@
{
"db_name": "PostgreSQL",
"query": "\n INSERT INTO auth.app_user (id, last_seen_at)\n VALUES ($1, now())\n ON CONFLICT (id)\n DO UPDATE SET last_seen_at = now()\n ",
"describe": {
"columns": [],
"parameters": {
"Left": [
"Uuid"
]
},
"nullable": []
},
"hash": "f389b37680c97698323c81cb7eeb9bcbafad49d7e1df83cb15a1b86bbdeb4fdf"
}
Generated
+1555 -33
View File
File diff suppressed because it is too large Load Diff
+22 -4
View File
@@ -3,12 +3,30 @@ name = "chat"
version = "0.1.0" version = "0.1.0"
edition = "2024" edition = "2024"
[dev-dependencies]
wiremock = "0.6"
tokio = { version = "1.52.3", features = ["macros", "rt-multi-thread"] }
[dependencies] [dependencies]
axum = "0.8.8" axum = { version = "0.8.9", features = ["macros"] }
utoipa = { version = "5.5.0", features = ["axum_extras", "uuid"] }
tokio = { version = "1", features = ["full"] } tokio = { version = "1", features = ["full"] }
serde = { version = "1", features = ["derive"] } serde = { version = "1", features = ["derive"] }
serde_json = "1" serde_json = "1.0.150"
jsonwebtoken = { version = "10.3.0", features = ["aws_lc_rs"] } jsonwebtoken = { version = "10.4.0", features = ["aws_lc_rs"] }
reqwest = { version = "0.13.2", features = ["json"] } reqwest = { version = "0.13.4", features = ["json", "stream"] }
once_cell = "1" once_cell = "1"
dotenvy = "0.15" dotenvy = "0.15"
thiserror = "2.0.18"
tokio-stream = "0.1"
futures = "0.3"
chrono = { version = "0.4.45", features = ["serde"] }
uuid = { version = "1.23.5", features = ["v4", "serde"] }
tower-http = { version = "0.7.0", features = ["cors"] }
tracing = "0.1.44"
tracing-subscriber = { version = "0.3", features = ["env-filter"]}
sqlx = { version = "0.8.6", features = ["runtime-tokio-rustls", "postgres", "uuid", "chrono", "macros"] }
rand = "0.8"
base64 = "0.22.1"
sha2 = "0.11.0"
async-stream = "0.3"
+1
View File
@@ -17,6 +17,7 @@ RUN rm -rf src
# Build actual app # Build actual app
COPY src ./src COPY src ./src
COPY .sqlx ./.sqlx
RUN cargo build --release RUN cargo build --release
# ── Runtime Stage ───────────────────────────────────────────────────────────── # ── Runtime Stage ─────────────────────────────────────────────────────────────
+7 -3
View File
@@ -7,12 +7,12 @@ services:
- .env - .env
ports: ports:
- "6066:3000" - "6066:3001"
restart: always restart: always
extra_hosts: networks:
- "host.docker.internal:host-gateway" - chat-net
environment: environment:
NVIDIA_DRIVER_CAPABILITIES: compute,utility,management NVIDIA_DRIVER_CAPABILITIES: compute,utility,management
@@ -23,3 +23,7 @@ services:
reservations: reservations:
devices: devices:
- capabilities: [gpu] - capabilities: [gpu]
networks:
chat-net:
external: true
+280 -1
View File
@@ -1,6 +1,285 @@
# TODO # TODO
- Race condition on jwks token refresh - Race condition on jwks token refresh
- Rate Limiting
git tag -d v1.0.0; git push origin :refs/tags/v1.0.0; git tag -a v1.0.0 -m "Release v1.0.0"; git push origin v1.0.0 git tag -d v1.0.0; git push origin :refs/tags/v1.0.0; git tag -a v1.0.0 -m "Release v1.0.0"; git push origin v1.0.0
curl https://chat.iceberg.black/api/v1/models -H "Authorization: Bearer eyJhbGciOiJSUzI1NiIsInR5cCIgOiAiSldUIiwia2lkIiA6ICJublpLek04TkZHVmpWbGFPRXZpMUtFSTVHQWRwaGlsYjh3RHRLeG5JOENZIn0.eyJleHAiOjE3NzU4MTUyODksImlhdCI6MTc3NTgxNDk4OSwianRpIjoiMTgwYzA2NDUtYzZiMC00MDRmLTgyYjEtMzA3YmY4ZmJlNmJlIiwiaXNzIjoiaHR0cHM6Ly9hdXRoLmljZWJlcmcuYmxhY2svcmVhbG1zL2ljZWJlcmciLCJzdWIiOiJmZGRiN2FjZC1kMmE5LTRmMTctOWIxNi1kZjVlN2EzNDI4YjciLCJ0eXAiOiJCZWFyZXIiLCJhenAiOiJjaGF0LWFwaSIsInNjb3BlIjoiIiwiY2xpZW50SG9zdCI6Ijg2LjIxMi44NC4xOTEiLCJjbGllbnRBZGRyZXNzIjoiODYuMjEyLjg0LjE5MSIsImNsaWVudF9pZCI6ImNoYXQtYXBpIn0.abjHABcjCJNiB6vriRw60nfzabEfD7CXwyRhkahFkC8ATgfy4fn0T8PnFfsaRpsXuhamIhWwNskrA7L9V3mbWWKW-JEOvImDhTc8sIX0E5fDTbk8O5wa_2yzNdLpRxdSjLqgL544rB8I-LZ8bl5SxtdN3gHfrnWr5ef8bbLgPRzZIylT3QUpah0uywDM_cfrrve9SMHHOUUItyzmOHLw0Igit1EzyFNyjbWf6OAU6TMjOF_eFTc5sakyBwsdJnGy0nhj5R-wxLpr1ug3iEd3Y-jDzHO4m6jariXVJ7Vvbz71i7sadDEKdSyOKcCeQbU4T0Tf7IxqW4c2DNnk3BFFuw"
curl -X POST "https://auth.iceberg.black/realms/iceberg/protocol/openid-connect/token" -H "Content-Type: application/x-www-form-urlencoded" -d "grant_type=client_credentials" -d "client_id=chat-api" -d "client_secret=5fHUp8Z5GoNM70MVOGuQTfKFkaAE48Za"
curl -s -X POST https://chat.iceberg.black/api/v1/chat/completions \
-H "Content-Type: application/json" \
-H "Authorization: Bearer eyJhbGciOiJSUzI1NiIsInR5cCIgOiAiSldUIiwia2lkIiA6ICJublpLek04TkZHVmpWbGFPRXZpMUtFSTVHQWRwaGlsYjh3RHRLeG5JOENZIn0.eyJleHAiOjE3NzU4MTU0NTEsImlhdCI6MTc3NTgxNTE1MSwianRpIjoiOGNjMmM5MjUtZmMzNy00MDFhLWE1ZjUtODFlOGJhNjcxNzcwIiwiaXNzIjoiaHR0cHM6Ly9hdXRoLmljZWJlcmcuYmxhY2svcmVhbG1zL2ljZWJlcmciLCJzdWIiOiJmZGRiN2FjZC1kMmE5LTRmMTctOWIxNi1kZjVlN2EzNDI4YjciLCJ0eXAiOiJCZWFyZXIiLCJhenAiOiJjaGF0LWFwaSIsInNjb3BlIjoiIiwiY2xpZW50SG9zdCI6Ijg2LjIxMi44NC4xOTEiLCJjbGllbnRBZGRyZXNzIjoiODYuMjEyLjg0LjE5MSIsImNsaWVudF9pZCI6ImNoYXQtYXBpIn0.Exv34aoAqwzSqQziZH3zScAnK_xiNCXqpOQl54x7NFMdO0gVgacJgMxoW82Ym8aCNfBRMFtWclZJ9RFh1b9uSEUgsdVtqvMX-2kCAhiMSjmIBPU-L0gr5N63c3bScOoAYa37baR4mXQ5LxHjfkLo_rHDJ74adg2JMo359Wbuu1_OR708_q8yjgO_4fQeYbIffxADwRDvPSIwQ8Y2qRjaAIOZs0xmA9p128CUxlUyvhxwfFquYKRaDs5QE9pIAtvm_KuoBydYopm8j8cbm4ixDhUwYjkzniIIapY77NxpmMosfh8BhK0W3ieK2gTKfmiFDhYgkGybU6b1BnJMy1_Kxg" \
-d '{
"model": "llama3:latest",
"messages": [
{"role": "user", "content": "What is Rust?"}
]
}' | jq .
curl -X POST http://localhost:3001/v1/completions \
-H "Content-Type: application/json" \
-H "x-api-key: UXVqi1Jazl_-A0TuBudw2Y3PeUiNwCMYayXwBWwuMf0" \
-H "Accept: text/event-stream" \
-d '{"model": "llama3:latest", "prompt": "hello", "stream": true}'
providers → LlmError ┐
├── services → ServiceError → api → HTTP
databases → DbError ┘
# 🦙 Ollama Rust API Wrapper
A high-performance Rust API wrapper around Ollama, providing an OpenAI-compatible interface, model lifecycle management, and advanced runtime features.
---
# 🚀 Features
- ✅ OpenAI-compatible API (`/v1/...`)
- ⚡ Streaming (Server-Sent Events)
- 🧠 Model lifecycle management (load/unload)
- 🔐 API key authentication (optional)
- 📊 Usage tracking & observability
- 🔀 Model routing & abstraction
- 🧩 Extensible architecture (multi-provider ready)
---
# 📡 API Endpoints
## 1. Core LLM API (OpenAI-compatible)
- [x] `POST /v1/chat/completions` + streaming
- [x] `POST /v1/completions` + streaming
- [ ] `POST /v1/embeddings`
- [x] `GET /v1/models`
## 2. Model Lifecycle Management
- [x] `POST /v1/models/{model}/load`
- [x] `POST /v1/models/{model}/unload`
## 3. Model Management
- [ ] `POST /v1/models/pull`
- [ ] `DELETE /v1/models/{model}`
## 4. Runtime & Observability
### Model Status
```
GET /v1/models/{model}/status
```
### List Loaded Models
```
GET /v1/runtime/models
```
---
## 6. Health Checks
```
GET /health
GET /ready
```
---
# 🧠 Internal Mapping (Ollama)
| Wrapper Endpoint | Ollama Endpoint |
| ------------------------- | --------------- |
| /v1/chat/completions | /api/chat |
| /v1/completions | /api/generate |
| /v1/embeddings | /api/embeddings |
| /v1/models | /api/tags |
| /v1/models/pull | /api/pull |
| DELETE /v1/models/{model} | /api/delete |
| load/unload | /api/generate |
---
# 🔧 Advanced Features
## 🔀 Model Routing
Use abstract model names:
```json
{
"model": "fast"
}
```
Example mapping:
```
fast → llama3:8b
smart → llama3:70b
code → deepseek-coder
```
---
## 📊 Usage Tracking
```
GET /v1/usage
```
Tracks:
- request count
- latency
- per-model usage
---
## 🚦 Rate Limiting
- Requests per minute
- Tokens per minute
Returns:
```
429 Too Many Requests
```
---
## 🧠 Sessions (Context Management)
```
POST /v1/sessions
POST /v1/sessions/{id}/chat
```
Stores conversation history server-side.
---
## ⚡ Caching
- Embeddings
- Deterministic prompts (temperature = 0)
---
## 🧩 Tool / Function Calling
Supports structured tool execution:
```json
{
"tools": [
{
"name": "function_name",
"parameters": {}
}
]
}
```
---
## 📦 Batch Requests
```
POST /v1/batch
```
---
## 🧠 Auto Eviction
```
POST /v1/runtime/evict
```
Strategies:
- LRU
- memory threshold
---
## 🧾 Logs
```
GET /v1/logs
```
## 🔔 Async Jobs / Webhooks
```
POST /v1/jobs
```
---
# 🏗️ Architecture
```
Client → Rust API → Ollama → Response
```
### Layers:
- HTTP (Axum)
- Service layer (business logic)
- Provider abstraction
- Ollama client
---
# 🔌 Provider Abstraction (Future-Proof)
```rust
trait LlmProvider {
async fn chat(...);
async fn embeddings(...);
}
```
Supports:
- Ollama (current)
- OpenAI (future)
- Others
---
# 🎯 Roadmap
- [ ] Full OpenAI compatibility
- [ ] Multi-node routing
- [ ] GPU-aware scheduling
- [ ] Web UI dashboard
- [ ] Distributed inference
---
# 🧠 Summary
This project turns Ollama into:
👉 A local OpenAI-compatible API
👉 A controllable model runtime
👉 A foundation for a full LLM gateway
# Bugs
## Ollama
- When model answer onoly with 1 tokens, answer is empty
# TODO
- open api doc for bearer token
- Unify check before sending to ollama payload
- load/unload model functions
+57
View File
@@ -0,0 +1,57 @@
use axum::{
Router,
http::{HeaderName, HeaderValue, Method, header},
};
use std::{env, sync::Arc};
use tower_http::cors::CorsLayer;
use crate::services::{AuthService, ChatService, ConversationService};
use crate::{api, databases::postgres, providers::ollama::client::OllamaProvider};
use once_cell::sync::Lazy;
static OLLAMA_URL: Lazy<String> = Lazy::new(|| env::var("OLLAMA_URL").expect("OLLAMA_URL not set"));
pub async fn build_app() -> Router {
let database_url = env::var("DATABASE_URL").expect("DATABASE_URL must be set");
let pool = postgres::pool::create_pool(&database_url)
.await
.expect("Fatal error");
let conversation_service = ConversationService::new(pool.clone());
let auth_service = AuthService::new(pool.clone());
let ollama = OllamaProvider::new(OLLAMA_URL.as_str());
let chat_service = ChatService::new(ollama, conversation_service.clone());
let state = Arc::new(api::state::AppState {
conversation_service,
auth_service,
chat_service,
});
let cors_origin =
env::var("CORS_ORIGIN").unwrap_or_else(|_| "http://localhost:3000".to_string());
let cors = CorsLayer::new()
.allow_origin(cors_origin.parse::<HeaderValue>().unwrap())
.allow_methods([
Method::GET,
Method::POST,
Method::PUT,
Method::DELETE,
Method::OPTIONS,
])
.allow_headers([
header::CONTENT_TYPE,
header::AUTHORIZATION,
header::ACCEPT,
HeaderName::from_static("x-api-key"),
])
.allow_credentials(true);
Router::new()
.nest("/v1", api::routes::v1::router(state.clone()))
.layer(cors)
.with_state(state)
}
+61
View File
@@ -0,0 +1,61 @@
use utoipa::OpenApi;
use crate::api;
use crate::api::routes;
#[derive(OpenApi)]
#[openapi(
info(
title = "Ollama Proxy",
description = "OpenAI-compatible proxy for local Ollama models",
version = "0.1.0",
license(
name = "MIT",
url = "https://opensource.org/licenses/MIT"
),
),
paths(
routes::v1::llm::list_models,
routes::v1::llm::completions,
routes::v1::llm::chat_completions,
routes::v1::llm::load_model,
routes::v1::llm::unload_model,
),
components(
schemas(
api::errors::ErrorResponse,
api::types::ApiModelsResponse,
api::types::ApiModelInfo,
api::types::ApiModelMetadata,
api::types::ApiLoadModelResponse,
api::types::ApiLoadModelRequest,
api::types::ApiUnloadModelResponse,
api::types::ApiLlmOptions,
api::types::ApiCompletionRequest,
api::types::ApiCompletionResponse,
api::types::ApiCompletionObject,
api::types::Choice,
api::types::Usage,
api::types::ApiFinishReason,
api::types::ApiChatRequest,
api::types::ApiMessage,
api::types::ApiRole,
api::types::ApiChatResponse,
api::types::ApiChatChoice,
api::types::ChatCompletionChunk,
api::types::ChatChunkChoice,
api::types::Delta,
)
),
tags(
(name = "chat", description = "Chat & completions"),
(name = "models", description = "Model management")
)
)]
pub struct ApiDoc;
+144
View File
@@ -0,0 +1,144 @@
use crate::api::middlewares::errors::AuthMiddlewareError;
use crate::databases::errors::DbError;
use crate::providers::keycloak::errors::JwtValidationError;
use crate::providers::ollama::errors::LlmError;
use crate::services::errors::ServiceError;
use axum::Json;
use axum::http::StatusCode;
use axum::response::{IntoResponse, Response};
use serde::Serialize;
use utoipa::ToSchema;
#[derive(Serialize, ToSchema)]
pub struct ErrorResponse {
pub status: u16,
pub error: String,
pub code: String,
}
#[derive(Debug, thiserror::Error)]
#[error("{message}")]
pub struct ApiError {
pub status: StatusCode,
pub code: &'static str,
pub message: String,
}
impl From<AuthMiddlewareError> for ApiError {
fn from(err: AuthMiddlewareError) -> Self {
match err {
AuthMiddlewareError::InvalidAuthorizationFormat => Self {
status: StatusCode::UNAUTHORIZED,
code: "AUTH_INVALID_FORMAT",
message: "invalid authorization format".into(),
},
AuthMiddlewareError::AuthenticationRequired => Self {
status: StatusCode::UNAUTHORIZED,
code: "AUTH_REQUIRED",
message: "missing authorization header or api key".into(),
},
AuthMiddlewareError::Forbidden => Self {
status: StatusCode::FORBIDDEN,
code: "AUTH_FORBIDDEN",
message: "insufficient permissions".into(),
},
}
}
}
impl From<ServiceError> for ApiError {
fn from(err: ServiceError) -> Self {
match err {
ServiceError::Db(e) => match e {
DbError::Connection(_) => Self {
status: StatusCode::INTERNAL_SERVER_ERROR,
code: "DB_CONNECTION",
message: "database connection error".into(),
},
DbError::Timeout => Self {
status: StatusCode::REQUEST_TIMEOUT,
code: "DB_TIMEOUT",
message: "database timeout".into(),
},
DbError::NotFound => Self {
status: StatusCode::NOT_FOUND,
code: "DB_NOT_FOUND",
message: "not found".into(),
},
DbError::Unauthorized => Self {
status: StatusCode::UNAUTHORIZED,
code: "DB_UNAUTHORIZED",
message: "unauthorized".into(),
},
},
ServiceError::Llm(e) => match e {
LlmError::MissingPrompt => Self {
status: StatusCode::BAD_REQUEST,
code: "OLLAMA_MISSING_PROMPT",
message: "prompt is required".into(),
},
LlmError::Http(_) => Self {
status: StatusCode::BAD_GATEWAY,
code: "OLLAMA_HTTP_ERROR",
message: "upstream error".into(),
},
_ => Self {
status: StatusCode::BAD_REQUEST,
code: "OLLAMA_ERROR",
message: "llm error".into(),
},
},
ServiceError::Auth(e) => match e {
JwtValidationError::InvalidToken | JwtValidationError::TokenValidationFailed => {
Self {
status: StatusCode::UNAUTHORIZED,
code: "AUTH_INVALID_TOKEN",
message: "invalid or expired token".into(),
}
}
JwtValidationError::InvalidHeader | JwtValidationError::MissingKid => Self {
status: StatusCode::UNAUTHORIZED,
code: "AUTH_INVALID_HEADER",
message: "invalid authorization header".into(),
},
JwtValidationError::JwkNotFound
| JwtValidationError::InvalidJwks
| JwtValidationError::MissingModulus
| JwtValidationError::MissingExponent
| JwtValidationError::InvalidDecodingKey => Self {
status: StatusCode::INTERNAL_SERVER_ERROR,
code: "AUTH_JWKS_ERROR",
message: "key validation error".into(),
},
JwtValidationError::JwksFetchFailed
| JwtValidationError::JwksRefreshFailed
| JwtValidationError::Reqwest(_) => Self {
status: StatusCode::SERVICE_UNAVAILABLE,
code: "AUTH_JWKS_FETCH",
message: "failed to fetch authorization keys".into(),
},
},
ServiceError::Internal(e) => Self {
status: StatusCode::INTERNAL_SERVER_ERROR,
code: "INTERNAL_SERVER_ERROR",
message: e,
},
}
}
}
impl IntoResponse for ApiError {
fn into_response(self) -> Response {
let body = ErrorResponse {
status: self.status.as_u16(),
error: self.message,
code: self.code.to_string(),
};
(self.status, Json(body)).into_response()
}
}
+178
View File
@@ -0,0 +1,178 @@
use crate::api::errors::ApiError;
use crate::api::middlewares::errors::AuthMiddlewareError;
use crate::api::state::SharedState;
use axum::{
extract::{Request, State},
http::{HeaderMap, StatusCode},
middleware::Next,
response::{IntoResponse, Response},
};
enum JwtError {
MissingHeader,
InvalidFormat,
Service(ApiError),
}
enum ApiKeyError {
MissingHeader,
Service(ApiError),
}
fn resolve_auth_error(jwt_err: JwtError, api_key_err: ApiKeyError) -> ApiError {
match (jwt_err, api_key_err) {
// Both headers absent
(JwtError::MissingHeader, ApiKeyError::MissingHeader) => {
tracing::debug!("auth failed: no credentials provided");
AuthMiddlewareError::AuthenticationRequired.into()
}
// JWT header present but malformed — api key result irrelevant
(JwtError::InvalidFormat, _) => {
tracing::debug!("auth failed: malformed Authorization header");
AuthMiddlewareError::InvalidAuthorizationFormat.into()
}
// JWT service failure, api key not attempted
(JwtError::Service(e), ApiKeyError::MissingHeader) => {
tracing::warn!(error = ?e, "auth failed: jwt validation error");
e
}
// Both services failed
(JwtError::Service(jwt_e), ApiKeyError::Service(api_e)) => {
tracing::warn!(
jwt_error = ?jwt_e,
api_key_error = ?api_e,
"auth failed: both jwt and api key validation errored"
);
jwt_e
}
// JWT missing, api key service failed
(JwtError::MissingHeader, ApiKeyError::Service(e)) => {
tracing::warn!(error = ?e, "auth failed: api key validation error");
e
}
}
}
pub async fn auth_middleware(
State(state): State<SharedState>,
mut req: Request,
next: Next,
) -> Response {
let headers = req.headers();
let auth = match try_jwt(&state, headers).await {
Ok(auth) => auth,
Err(jwt_err) => match try_api_key(&state, headers).await {
Ok(auth) => auth,
Err(api_key_err) => {
return resolve_auth_error(jwt_err, api_key_err).into_response();
}
},
};
if let Err(err) = record_auth(&state, &auth).await {
tracing::debug!("Error during authentification {:?}", err);
return err.into_response();
}
tracing::debug!("User authentified");
req.extensions_mut().insert(auth);
next.run(req).await
}
async fn try_jwt(
state: &SharedState,
headers: &HeaderMap,
) -> Result<crate::core::auth::Auth, JwtError> {
println!("{:?}", headers);
let token = headers
.get("authorization")
.and_then(|v| v.to_str().ok())
.ok_or(JwtError::MissingHeader)?
.strip_prefix("Bearer ")
.ok_or(JwtError::InvalidFormat)?;
let claims = state
.auth_service
.validate_jwt(token)
.await
.map_err(|e| JwtError::Service(e.into()))?;
Ok(crate::core::auth::Auth::Jwt(claims))
}
async fn try_api_key(
state: &SharedState,
headers: &HeaderMap,
) -> Result<crate::core::auth::Auth, ApiKeyError> {
let key = headers
.get("x-api-key")
.and_then(|v| v.to_str().ok())
.ok_or(ApiKeyError::MissingHeader)?;
let auth = state
.auth_service
.validate_api_key(key)
.await
.map_err(|e| ApiKeyError::Service(e.into()))?;
Ok(crate::core::auth::Auth::ApiKey(auth))
}
async fn record_auth(state: &SharedState, auth: &crate::core::auth::Auth) -> Result<(), ApiError> {
match auth {
crate::core::auth::Auth::Jwt(_) => {
state.auth_service.create_user(&auth.user_id()).await?;
}
crate::core::auth::Auth::ApiKey(key) => {
state
.auth_service
.update_last_access_api_key(&key.user_id)
.await?;
}
}
Ok(())
}
// ── Role guard ───────────────────────────────────────────────────────────────
#[macro_export]
macro_rules! role_guard {
($jwt:expr, $api:expr) => {
middleware::from_fn(move |req, next| {
$crate::api::middlewares::auth::require_roles(req, next, $jwt, $api)
})
};
}
pub async fn require_roles(
request: Request,
next: Next,
jwt_role: Option<&'static str>,
api_key_role: Option<&crate::core::auth::api_key::KeyRole>,
) -> Result<Response, ApiError> {
let auth = request
.extensions()
.get::<crate::core::auth::Auth>()
.ok_or(ApiError {
status: StatusCode::UNAUTHORIZED,
code: "MISSING_AUTH_HEADER",
message: "API key missing".into(),
})?;
if jwt_role.is_some_and(|role| !auth.has_jwt_role(role)) {
return Err(AuthMiddlewareError::Forbidden.into());
}
if api_key_role.is_some_and(|role| !auth.has_apikey_role(role)) {
return Err(AuthMiddlewareError::Forbidden.into());
}
Ok(next.run(request).await)
}
+13
View File
@@ -0,0 +1,13 @@
use thiserror::Error;
#[derive(Debug, Error)]
pub enum AuthMiddlewareError {
#[error("invalid authorization format")]
InvalidAuthorizationFormat,
#[error("authentication required")]
AuthenticationRequired,
#[error("insufficient permissions")]
Forbidden,
}
+2
View File
@@ -0,0 +1,2 @@
pub mod auth;
pub mod errors;
+7
View File
@@ -0,0 +1,7 @@
pub mod app;
pub mod docs;
pub mod errors;
pub mod middlewares;
pub mod routes;
pub mod state;
pub mod types;
+1
View File
@@ -0,0 +1 @@
pub mod v1;
+25
View File
@@ -0,0 +1,25 @@
use crate::api::errors::ApiError;
use crate::api::state::SharedState;
use crate::api::types::{CreateApiKeyRequest, CreateApiKeyResponse};
use crate::core::auth::Auth;
use axum::{
Json,
extract::{Extension, State},
};
pub async fn create_api_key(
State(state): State<SharedState>,
Extension(claims): Extension<Auth>,
Json(body): Json<CreateApiKeyRequest>,
) -> Result<Json<CreateApiKeyResponse>, ApiError> {
let payload = crate::core::auth::api_key::CreateApiKeyRequest {
user_id: claims.user_id(),
name: body.name,
roles: body.scopes.into_iter().map(Into::into).collect(),
};
let key = state.auth_service.create_api_key(payload).await?;
Ok(Json(CreateApiKeyResponse { api_key: key }))
}
+384
View File
@@ -0,0 +1,384 @@
use crate::api;
use crate::api::errors::{ApiError, ErrorResponse};
use crate::api::state::SharedState;
use crate::core::auth::Auth;
use crate::core::llm::completions::CompletionResult;
use axum::{
Json,
extract::{Extension, Path, Query, State},
response::{
IntoResponse,
sse::{Event, KeepAlive, Sse},
},
};
use futures::StreamExt;
use uuid::Uuid;
#[utoipa::path(
get,
path = "/llm/models",
tag = "models",
responses(
(
status = 200,
description = "List of locally available Ollama models",
body = api::types::ApiModelsResponse,
content_type = "application/json",
),
(
status = 500,
description = "Internal server error (Ollama or network failure)",
body = api::errors::ErrorResponse,
example = json!({ "error": "connection refused" })
)
)
)]
pub async fn list_models(
State(state): State<SharedState>,
) -> Result<Json<api::types::ApiModelsResponse>, api::errors::ApiError> {
let models = state.chat_service.list_models().await?;
Ok(Json(models.into()))
}
#[utoipa::path(
post,
path = "/llm/models/{model}/load",
tag = "models",
params(
("model" = String, Path, description = "Name of the model to load into memory (e.g. 'llama3')")
),
request_body(
content = api::types::ApiLoadModelRequest,
description = "Load model request",
content_type = "application/json",
example = json!({ "keep_alive": "10m" })
),
responses(
(
status = 200,
description = "Model successfully loaded into memory",
body = api::types::ApiLoadModelResponse,
content_type = "application/json",
),
(
status = 400,
description = "Invalid or missing keep_alive format",
body = api::errors::ErrorResponse,
examples(
("Missing" = (value = json!({ "error": "keep alive is required and cannot be empty" }))),
("Invalid" = (value = json!({ "error": "invalid keep_alive '10x' — use 30s / 10m / 2h, a plain integer, or -1" })))
)
),
(
status = 404,
description = "Model not found locally",
body = api::errors::ErrorResponse,
example = json!({ "error": "model 'llama3' not found — run `ollama pull llama3`" })
),
(
status = 500,
description = "Internal server error (Ollama or network failure)",
body = api::errors::ErrorResponse,
example = json!({ "error": "connection refused" })
)
)
)]
pub async fn load_model(
State(state): State<SharedState>,
Path(model): Path<String>,
Json(body): Json<api::types::ApiLoadModelRequest>,
) -> Result<Json<api::types::ApiLoadModelResponse>, api::errors::ApiError> {
let response = state
.chat_service
.load_model(crate::core::llm::models::LoadModelRequest {
model,
keep_alive: body.keep_alive.clone(),
})
.await?;
Ok(Json(api::types::ApiLoadModelResponse {
model: response.model,
keep_alive: body.keep_alive,
status: "loaded".to_string(),
}))
}
#[utoipa::path(
post,
path = "/llm/models/{model}/unload",
tag = "models",
params(
("model" = String, Path, description = "Name of the model to unload from memory (e.g. 'llama3')")
),
responses(
(
status = 200,
description = "Model successfully unloaded from memory",
body = api::types::ApiUnloadModelResponse,
content_type = "application/json",
),
(
status = 404,
description = "Model not found locally",
body = api::errors::ErrorResponse,
example = json!({ "error": "model 'llama3' not found — run `ollama pull llama3`" })
),
(
status = 500,
description = "Internal server error (Ollama or network failure)",
body = api::errors::ErrorResponse,
example = json!({ "error": "connection refused" })
)
)
)]
pub async fn unload_model(
State(state): State<SharedState>,
Path(model): Path<String>,
) -> Result<Json<api::types::ApiUnloadModelResponse>, api::errors::ApiError> {
let response = state
.chat_service
.unload_model(crate::core::llm::models::UnloadModelRequest { model })
.await?;
Ok(Json(api::types::ApiUnloadModelResponse {
model: response.model,
status: "unloaded".to_string(),
}))
}
#[utoipa::path(
post,
path = "/llm/completions",
tag = "chat",
request_body(
content = api::types::ApiCompletionRequest,
description = "Text completion request",
content_type = "application/json"
),
responses(
(
status = 200,
description = "Text completion response. If stream=true, response is SSE stream of chunks ending in [DONE].",
body = api::types::ApiCompletionResponse,
content_type = "application/json"
),
(
status = 400,
description = "Invalid request: missing prompt, model, or invalid format",
body = ErrorResponse,
example = json!({ "error": "prompt is required and cannot be empty" })
),
(
status = 422,
description = "Model not found or not available locally",
body = ErrorResponse,
example = json!({ "error": "model 'llama3' is not available — run `ollama pull llama3` first" })
),
(
status = 500,
description = "Internal server error (Ollama or network failure)",
body = ErrorResponse,
example = json!({ "error": "connection refused" })
)
)
)]
pub async fn completions(
State(state): State<SharedState>,
Json(body): Json<api::types::ApiCompletionRequest>,
) -> Result<impl IntoResponse, ApiError> {
tracing::debug!("Received /completion with body {:?}", body);
let response = state.chat_service.complete(body.into()).await?;
match response {
CompletionResult::NoStream(res) => {
Ok(Json::<crate::api::types::ApiCompletionResponse>(res.into()).into_response())
}
CompletionResult::Stream(stream) => {
let sse_stream = stream.map(|item| match item {
Ok(event) => {
let data = serde_json::to_string(&event).unwrap_or_default();
Ok(Event::default().data(data))
}
Err(e) => Err(e),
});
Ok(Sse::new(sse_stream)
.keep_alive(KeepAlive::default())
.into_response())
}
}
}
#[utoipa::path(
post,
path = "/llm/chat/completions",
tag = "chat",
request_body(
content = api::types::ApiChatRequest,
description = "Chat completion request with message history",
content_type = "application/json"
),
responses(
(
status = 200,
description = "Chat completion response. If stream=false returns JSON. If stream=true returns SSE stream of chunks ending with [DONE].",
body = api::types::ApiChatResponse,
content_type = "application/json"
),
(
status = 400,
description = "Invalid request",
body = api::errors::ErrorResponse,
example = json!({ "error": "messages array with at least one user message is required" })
),
(
status = 401,
description = "Unauthorized",
body = api::errors::ErrorResponse,
example = json!({ "error": "missing or invalid token" })
),
(
status = 422,
description = "Model not found or unavailable",
body = api::errors::ErrorResponse,
example = json!({ "error": "model 'llama3' is not available — run `ollama pull llama3` first" })
),
(
status = 500,
description = "Internal server error",
body = api::errors::ErrorResponse,
example = json!({ "error": "connection refused" })
)
)
)]
pub async fn chat_completions(
State(state): State<SharedState>,
Extension(auth): Extension<Auth>,
Json(body): Json<api::types::ApiChatRequest>,
) -> Result<impl IntoResponse, ApiError> {
tracing::debug!("Received /chat/completion with body {:?}", body);
let response = state.chat_service.chat_complete(body.into(), &auth).await?;
match response {
crate::core::llm::chat::ChatCompletionResult::NoStream(res) => {
Ok(Json::<crate::api::types::ApiChatResponse>(res.into()).into_response())
}
crate::core::llm::chat::ChatCompletionResult::Stream(stream) => {
let sse_stream = stream.map(|item| -> Result<Event, ApiError> {
tracing::debug!("{:?}", item);
match item {
Err(e) => {
tracing::debug!("Error in chat_completions: {:?}", e);
Err(ApiError::from(e))
}
Ok(crate::core::llm::chat::ChatCompletionStreamEvent::Start {
conversation_id,
message_id,
created_at,
model,
}) => {
println!("RECEIVED STRATTT");
let payload = serde_json::to_string(&api::types::StreamEvent::Start(
api::types::StartEventData {
conversation_id,
created: created_at,
id: message_id,
model,
},
))
.unwrap_or_default();
Ok(Event::default().data(payload))
}
Ok(crate::core::llm::chat::ChatCompletionStreamEvent::Token {
content,
id,
..
}) => {
let chunk = api::types::ChatCompletionChunk {
id,
object: "chat.completion.chunk".to_string(),
choices: vec![api::types::ChatChunkChoice {
index: 0,
delta: api::types::Delta {
content: Some(content),
role: Some(api::types::ApiRole::Assistant),
},
}],
};
let payload = serde_json::to_string(&api::types::StreamEvent::Delta(chunk))
.unwrap_or_default();
Ok(Event::default().data(payload))
}
Ok(crate::core::llm::chat::ChatCompletionStreamEvent::Final(res)) => {
let payload = serde_json::to_string(&api::types::StreamEvent::End(
api::types::EndEventData {
id: res.id,
usage: api::types::Usage {
prompt_tokens: res.prompt_tokens,
completion_tokens: res.completion_tokens,
total_tokens: res.prompt_tokens + res.completion_tokens,
},
finish_reason: res.finish_reason.into(),
},
))
.unwrap_or_default();
Ok(Event::default().data(payload))
}
}
});
Ok(Sse::new(sse_stream)
.keep_alive(KeepAlive::default())
.into_response())
}
}
}
// Conversation retrieveing
pub async fn get_conversations(
State(state): State<SharedState>,
Extension(auth): Extension<Auth>,
Query(params): Query<api::types::CursorPage>,
) -> Result<Json<api::types::ConversationListResponse>, ApiError> {
tracing::debug!("Conversation hit: {:?}", auth);
let pointer = crate::core::databases::conversations::CursorPage {
limit: params.limit.unwrap_or(10),
before: params.before,
};
let conversations = state
.conversation_service
.get_conversations_entries(auth.user_id(), pointer)
.await?;
Ok(Json(conversations.into()))
}
pub async fn get_messages(
State(state): State<SharedState>,
Extension(auth): Extension<Auth>,
Path(conversation_id): Path<Uuid>,
Query(params): Query<api::types::CursorPage>,
) -> Result<Json<api::types::MessageListResponse>, ApiError> {
tracing::debug!("Messages hit: {:?}", auth);
let pointer = crate::core::databases::conversations::CursorPage {
limit: params.limit.unwrap_or(15),
before: params.before,
};
let messages = state
.conversation_service
.get_messages_entries(auth.user_id(), conversation_id, pointer)
.await?;
Ok(Json(messages.into()))
}
+58
View File
@@ -0,0 +1,58 @@
pub mod apikey;
pub mod llm;
use crate::api::docs::ApiDoc;
use crate::api::middlewares::auth::auth_middleware;
use crate::api::state::app_state::SharedState;
use crate::role_guard;
use axum::{
Json, Router, middleware,
routing::{get, post},
};
use utoipa::OpenApi;
async fn openapi_json() -> Json<utoipa::openapi::OpenApi> {
Json(ApiDoc::openapi())
}
fn llm_router() -> Router<SharedState> {
Router::new()
.route("/models", get(llm::list_models))
.route("/completions", post(llm::completions))
.route("/chat/completions", post(llm::chat_completions))
.route("/models/{model}/load", post(llm::load_model))
.route("/models/{model}/unload", post(llm::unload_model))
}
fn keys_router() -> Router<SharedState> {
Router::new().route(
"/generate",
post(apikey::create_api_key).route_layer(role_guard!(Some("admin"), None)),
)
}
fn log_router() -> Router<SharedState> {
Router::new()
.route("/conversations", get(llm::get_conversations))
.route(
"/conversations/{conversation_id}/messages",
get(llm::get_messages),
)
}
pub fn router(state: SharedState) -> Router<SharedState> {
let protected = Router::new()
.nest("/llm", llm_router())
.nest("/keys", keys_router())
.nest("/log", log_router())
.route_layer(middleware::from_fn_with_state(
state.clone(),
auth_middleware,
));
Router::new()
.route("/docs.json", get(openapi_json))
.merge(protected)
.with_state(state)
}
+574
View File
@@ -0,0 +1,574 @@
use axum::{
body::Body,
extract::{Multipart, Path, Query, State},
http::header,
response::{IntoResponse, Response},
Json,
};
use serde::{Deserialize, Serialize};
use utoipa::ToSchema;
use uuid::Uuid;
use crate::api;
use crate::api::errors::ApiError;
use crate::SharedState;
// ---------------------------------------------------------------------
// Types
// ---------------------------------------------------------------------
#[derive(Debug, Clone, Copy, Serialize, Deserialize, ToSchema)]
#[serde(rename_all = "lowercase")]
pub enum AudioFormat {
Mp3,
Wav,
Ogg,
Pcm,
}
impl AudioFormat {
pub fn content_type(&self) -> &'static str {
match self {
AudioFormat::Mp3 => "audio/mpeg",
AudioFormat::Wav => "audio/wav",
AudioFormat::Ogg => "audio/ogg",
AudioFormat::Pcm => "audio/L16",
}
}
}
#[derive(Debug, Deserialize, ToSchema)]
pub struct ApiSpeechRequest {
/// Text to synthesize.
#[schema(example = "Hello there, how can I help you today?")]
pub text: String,
/// Voice id, as returned by `GET /v1/voices`.
#[schema(example = "voice_en_us_amy")]
pub voice: String,
/// TTS model/engine to use. Defaults to the server's default model.
#[schema(example = "xtts-v2")]
pub model: Option<String>,
/// BCP-47 language code. Defaults to the voice's native language.
#[schema(example = "en-US")]
pub language: Option<String>,
/// Output audio format. Defaults to mp3.
pub format: Option<AudioFormat>,
/// Playback speed multiplier (0.5-2.0). Defaults to 1.0.
#[schema(example = 1.0)]
pub speed: Option<f32>,
}
#[derive(Debug, Serialize, ToSchema)]
pub struct ApiVoice {
pub id: String,
pub name: String,
/// BCP-47 language code, e.g. "en-US".
pub language: String,
pub sample_rate: u32,
/// URL to a short preview clip, if available.
pub preview_url: Option<String>,
pub is_cloned: bool,
pub created_at: chrono::DateTime<chrono::Utc>,
}
#[derive(Debug, Serialize, ToSchema)]
pub struct ApiVoiceListResponse {
pub voices: Vec<ApiVoice>,
}
#[derive(Debug, Deserialize, ToSchema)]
pub struct ApiRegisterVoiceRequest {
/// Display name for the new voice.
#[schema(example = "My Cloned Voice")]
pub name: String,
/// BCP-47 language code for the voice.
#[schema(example = "en-US")]
pub language: Option<String>,
}
#[derive(Debug, Serialize, ToSchema)]
pub struct ApiRegisterVoiceResponse {
pub voice: ApiVoice,
}
#[derive(Debug, Serialize, ToSchema)]
pub struct ApiDeleteVoiceResponse {
pub id: String,
pub deleted: bool,
}
#[derive(Debug, Serialize, ToSchema)]
pub struct ApiLanguage {
/// BCP-47 language code, e.g. "en-US".
pub code: String,
pub name: String,
}
#[derive(Debug, Serialize, ToSchema)]
pub struct ApiLanguageListResponse {
pub languages: Vec<ApiLanguage>,
}
#[derive(Debug, Serialize, ToSchema)]
pub struct ApiTtsModel {
pub id: String,
pub name: String,
pub description: Option<String>,
pub supported_languages: Vec<String>,
}
#[derive(Debug, Serialize, ToSchema)]
pub struct ApiTtsModelListResponse {
pub models: Vec<ApiTtsModel>,
}
// ---------------------------------------------------------------------
// POST /v1/audio/speech
// ---------------------------------------------------------------------
#[utoipa::path(
post,
path = "/v1/audio/speech",
tag = "audio",
request_body(
content = ApiSpeechRequest,
description = "Speech synthesis request",
content_type = "application/json",
example = json!({ "text": "Hello there!", "voice": "voice_en_us_amy", "format": "mp3" })
),
responses(
(
status = 200,
description = "Full audio file generated from the input text",
content_type = "audio/mpeg",
),
(
status = 400,
description = "Invalid request (e.g. empty text, unsupported speed)",
body = api::errors::ErrorResponse,
example = json!({ "error": "text must not be empty" })
),
(
status = 404,
description = "Voice or model not found",
body = api::errors::ErrorResponse,
example = json!({ "error": "voice 'voice_en_us_amy' not found" })
),
(
status = 500,
description = "Internal server error during synthesis",
body = api::errors::ErrorResponse,
example = json!({ "error": "synthesis engine crashed" })
)
)
)]
pub async fn generate_speech(
State(state): State<SharedState>,
Json(body): Json<ApiSpeechRequest>,
) -> Result<Response, ApiError> {
let audio = state
.tts_service
.synthesize(crate::core::tts::SynthesizeRequest {
text: body.text,
voice: body.voice,
model: body.model,
language: body.language,
format: body.format.unwrap_or(AudioFormat::Mp3),
speed: body.speed.unwrap_or(1.0),
})
.await?;
let content_type = audio.format.content_type();
Ok((
[(header::CONTENT_TYPE, content_type)],
Body::from(audio.bytes),
)
.into_response())
}
// ---------------------------------------------------------------------
// POST /v1/audio/speech/stream
// ---------------------------------------------------------------------
#[utoipa::path(
post,
path = "/v1/audio/speech/stream",
tag = "audio",
request_body(
content = ApiSpeechRequest,
description = "Speech synthesis request, streamed back as audio is generated",
content_type = "application/json",
example = json!({ "text": "Hello there!", "voice": "voice_en_us_amy", "format": "mp3" })
),
responses(
(
status = 200,
description = "Chunked audio stream (chunk-transfer-encoded); the same audio the sync endpoint returns, sent incrementally",
content_type = "audio/mpeg",
),
(
status = 400,
description = "Invalid request",
body = api::errors::ErrorResponse,
example = json!({ "error": "text must not be empty" })
),
(
status = 404,
description = "Voice or model not found",
body = api::errors::ErrorResponse,
example = json!({ "error": "voice 'voice_en_us_amy' not found" })
),
(
status = 500,
description = "Internal server error during synthesis",
body = api::errors::ErrorResponse,
example = json!({ "error": "synthesis engine crashed" })
)
)
)]
pub async fn generate_speech_stream(
State(state): State<SharedState>,
Json(body): Json<ApiSpeechRequest>,
) -> Result<Response, ApiError> {
let format = body.format.unwrap_or(AudioFormat::Mp3);
let stream = state
.tts_service
.synthesize_stream(crate::core::tts::SynthesizeRequest {
text: body.text,
voice: body.voice,
model: body.model,
language: body.language,
format,
speed: body.speed.unwrap_or(1.0),
})
.await?;
Ok((
[(header::CONTENT_TYPE, format.content_type())],
Body::from_stream(stream),
)
.into_response())
}
// ---------------------------------------------------------------------
// GET /v1/voices
// ---------------------------------------------------------------------
#[derive(Debug, Deserialize, ToSchema)]
pub struct ListVoicesQuery {
/// Optional BCP-47 language filter, e.g. "en-US".
pub language: Option<String>,
}
#[utoipa::path(
get,
path = "/v1/voices",
tag = "voices",
params(
("language" = Option<String>, Query, description = "Filter voices by BCP-47 language code")
),
responses(
(
status = 200,
description = "List of available voices",
body = ApiVoiceListResponse,
content_type = "application/json",
),
(
status = 500,
description = "Internal server error",
body = api::errors::ErrorResponse,
example = json!({ "error": "failed to load voice registry" })
)
)
)]
pub async fn list_voices(
State(state): State<SharedState>,
Query(query): Query<ListVoicesQuery>,
) -> Result<Json<ApiVoiceListResponse>, ApiError> {
let voices = state.tts_service.list_voices(query.language).await?;
Ok(Json(ApiVoiceListResponse {
voices: voices.into_iter().map(Into::into).collect(),
}))
}
// ---------------------------------------------------------------------
// POST /v1/voices (register / clone a voice)
// ---------------------------------------------------------------------
#[utoipa::path(
post,
path = "/v1/voices",
tag = "voices",
request_body(
content = ApiRegisterVoiceRequest,
description = "Multipart form: JSON fields (name, language) plus an `audio_sample` file field containing the reference audio to clone",
content_type = "multipart/form-data",
),
responses(
(
status = 200,
description = "Voice registered/cloned successfully",
body = ApiRegisterVoiceResponse,
content_type = "application/json",
),
(
status = 400,
description = "Invalid request (missing name, missing/unsupported audio sample)",
body = api::errors::ErrorResponse,
examples(
("Missing name" = (value = json!({ "error": "name is required and cannot be empty" }))),
("Bad sample" = (value = json!({ "error": "audio_sample must be a wav or mp3 file under 30s" })))
)
),
(
status = 500,
description = "Internal server error while cloning the voice",
body = api::errors::ErrorResponse,
example = json!({ "error": "voice cloning engine failed" })
)
)
)]
pub async fn register_voice(
State(state): State<SharedState>,
mut multipart: Multipart,
) -> Result<Json<ApiRegisterVoiceResponse>, ApiError> {
let mut name: Option<String> = None;
let mut language: Option<String> = None;
let mut audio_sample: Option<Vec<u8>> = None;
while let Some(field) = multipart
.next_field()
.await
.map_err(|e| ApiError::bad_request(format!("invalid multipart body: {e}")))?
{
match field.name() {
Some("name") => {
name = Some(field.text().await.map_err(|e| {
ApiError::bad_request(format!("invalid name field: {e}"))
})?);
}
Some("language") => {
language = Some(field.text().await.map_err(|e| {
ApiError::bad_request(format!("invalid language field: {e}"))
})?);
}
Some("audio_sample") => {
audio_sample = Some(
field
.bytes()
.await
.map_err(|e| {
ApiError::bad_request(format!("invalid audio_sample field: {e}"))
})?
.to_vec(),
);
}
_ => {}
}
}
let name = name.ok_or_else(|| {
ApiError::bad_request("name is required and cannot be empty".to_string())
})?;
let audio_sample = audio_sample.ok_or_else(|| {
ApiError::bad_request("audio_sample is required".to_string())
})?;
let voice = state
.tts_service
.register_voice(crate::core::tts::RegisterVoiceRequest {
name,
language,
audio_sample,
})
.await?;
Ok(Json(ApiRegisterVoiceResponse {
voice: voice.into(),
}))
}
// ---------------------------------------------------------------------
// GET /v1/voices/{id}
// ---------------------------------------------------------------------
#[utoipa::path(
get,
path = "/v1/voices/{id}",
tag = "voices",
params(
("id" = String, Path, description = "Voice id, e.g. 'voice_en_us_amy'")
),
responses(
(
status = 200,
description = "Voice metadata (language, sample rate, preview URL)",
body = ApiVoice,
content_type = "application/json",
),
(
status = 404,
description = "Voice not found",
body = api::errors::ErrorResponse,
example = json!({ "error": "voice 'voice_en_us_amy' not found" })
),
(
status = 500,
description = "Internal server error",
body = api::errors::ErrorResponse,
example = json!({ "error": "failed to load voice registry" })
)
)
)]
pub async fn get_voice(
State(state): State<SharedState>,
Path(id): Path<String>,
) -> Result<Json<ApiVoice>, ApiError> {
let voice = state.tts_service.get_voice(&id).await?;
Ok(Json(voice.into()))
}
// ---------------------------------------------------------------------
// DELETE /v1/voices/{id}
// ---------------------------------------------------------------------
#[utoipa::path(
delete,
path = "/v1/voices/{id}",
tag = "voices",
params(
("id" = String, Path, description = "Voice id to remove")
),
responses(
(
status = 200,
description = "Voice removed successfully",
body = ApiDeleteVoiceResponse,
content_type = "application/json",
),
(
status = 404,
description = "Voice not found",
body = api::errors::ErrorResponse,
example = json!({ "error": "voice 'voice_en_us_amy' not found" })
),
(
status = 400,
description = "Voice is a built-in voice and cannot be deleted",
body = api::errors::ErrorResponse,
example = json!({ "error": "built-in voices cannot be deleted" })
),
(
status = 500,
description = "Internal server error",
body = api::errors::ErrorResponse,
example = json!({ "error": "failed to update voice registry" })
)
)
)]
pub async fn delete_voice(
State(state): State<SharedState>,
Path(id): Path<String>,
) -> Result<Json<ApiDeleteVoiceResponse>, ApiError> {
state.tts_service.delete_voice(&id).await?;
Ok(Json(ApiDeleteVoiceResponse {
id,
deleted: true,
}))
}
// ---------------------------------------------------------------------
// GET /v1/languages
// ---------------------------------------------------------------------
#[utoipa::path(
get,
path = "/v1/languages",
tag = "languages",
responses(
(
status = 200,
description = "List of supported languages",
body = ApiLanguageListResponse,
content_type = "application/json",
),
(
status = 500,
description = "Internal server error",
body = api::errors::ErrorResponse,
example = json!({ "error": "failed to load language list" })
)
)
)]
pub async fn list_languages(
State(state): State<SharedState>,
) -> Result<Json<ApiLanguageListResponse>, ApiError> {
let languages = state.tts_service.list_languages().await?;
Ok(Json(ApiLanguageListResponse {
languages: languages
.into_iter()
.map(|l| ApiLanguage {
code: l.code,
name: l.name,
})
.collect(),
}))
}
// ---------------------------------------------------------------------
// GET /v1/models
// ---------------------------------------------------------------------
#[utoipa::path(
get,
path = "/v1/models",
tag = "models",
responses(
(
status = 200,
description = "List of available TTS models/engines",
body = ApiTtsModelListResponse,
content_type = "application/json",
),
(
status = 500,
description = "Internal server error",
body = api::errors::ErrorResponse,
example = json!({ "error": "failed to enumerate models" })
)
)
)]
pub async fn list_models(
State(state): State<SharedState>,
) -> Result<Json<ApiTtsModelListResponse>, ApiError> {
let models = state.tts_service.list_models().await?;
Ok(Json(ApiTtsModelListResponse {
models: models
.into_iter()
.map(|m| ApiTtsModel {
id: m.id,
name: m.name,
description: m.description,
supported_languages: m.supported_languages,
})
.collect(),
}))
}
// ---------------------------------------------------------------------
// Router wiring (example)
// ---------------------------------------------------------------------
pub fn router() -> axum::Router<SharedState> {
use axum::routing::{delete, get, post};
axum::Router::new()
.route("/v1/audio/speech", post(generate_speech))
.route("/v1/audio/speech/stream", post(generate_speech_stream))
.route("/v1/voices", get(list_voices).post(register_voice))
.route("/v1/voices/{id}", get(get_voice).delete(delete_voice))
.route("/v1/languages", get(list_languages))
.route("/v1/models", get(list_models))
}
+12
View File
@@ -0,0 +1,12 @@
use crate::services::{AuthService, ChatService, ConversationService};
use std::sync::Arc;
#[derive(Clone)]
pub struct AppState {
pub conversation_service: ConversationService,
pub auth_service: AuthService,
pub chat_service: ChatService,
}
pub type SharedState = Arc<AppState>;
+4
View File
@@ -0,0 +1,4 @@
pub mod app_state;
pub use app_state::AppState;
pub use app_state::SharedState;
+283
View File
@@ -0,0 +1,283 @@
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use utoipa::ToSchema;
use uuid::Uuid;
// ------ Models ------
#[derive(Debug, Serialize, ToSchema)]
pub struct ApiModelInfo {
pub name: String,
pub family: Option<String>,
pub parameter_size: Option<String>,
pub metadata: ApiModelMetadata,
}
#[derive(Debug, Clone, Default, Serialize, ToSchema)]
pub struct ApiModelMetadata {
pub extra: HashMap<String, String>,
}
// ------ Endpoint: /models ------
#[derive(Debug, Serialize, ToSchema)]
pub struct ApiModelsResponse {
pub models: Vec<ApiModelInfo>,
}
// ------ Endpoint: /models/{model}/load ------
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct ApiLoadModelResponse {
pub model: String,
pub status: String,
pub keep_alive: String,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct ApiLoadModelRequest {
pub keep_alive: String,
}
// ------ Endpoint: /models/{model}/unload ------
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct ApiUnloadModelResponse {
pub model: String,
pub status: String,
}
// ------ Completions ------
#[derive(Debug, Deserialize, Serialize, ToSchema, Default)]
pub struct ApiLlmOptions {
#[serde(default)]
pub stream: bool,
pub temperature: Option<f32>,
pub top_p: Option<f32>,
pub top_k: Option<u32>,
pub repeat_penalty: Option<f32>,
pub seed: Option<i64>,
pub num_ctx: Option<u32>,
pub num_predict: Option<u32>,
pub stop: Option<Vec<String>>,
pub keep_alive: Option<String>,
pub context_depth: Option<u32>,
}
// ------ Endpoint: /completions ------
#[derive(Debug, Deserialize, Serialize, ToSchema)]
pub struct ApiCompletionRequest {
#[serde(flatten)]
pub options: ApiLlmOptions,
pub model: String,
pub prompt: String,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct ApiCompletionResponse {
pub id: Uuid,
pub object: ApiCompletionObject,
pub created: String,
pub model: String,
pub choices: Vec<Choice>,
pub usage: Usage,
}
#[derive(Debug, Serialize, Deserialize, ToSchema, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum ApiCompletionObject {
TextCompletion,
}
// ------ Endpoint: /chat/completions ------
#[derive(Debug, Serialize, Deserialize, ToSchema, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum ApiFinishReason {
Stop,
Length,
Error,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct Choice {
pub text: String,
pub index: u32,
pub finish_reason: ApiFinishReason,
}
#[derive(Debug, Serialize, Deserialize, ToSchema, Clone, Copy)]
pub struct Usage {
pub prompt_tokens: u32,
pub completion_tokens: u32,
pub total_tokens: u32,
}
#[derive(Debug, Deserialize, Serialize, ToSchema)]
pub struct ApiChatRequest {
#[serde(flatten)]
pub base: ApiLlmOptions,
pub model: String,
pub message: ApiMessage,
// Non standard Open AI
pub conversation_id: Option<Uuid>,
pub parent_id: Option<Uuid>, // used when branching, regenerate, etc
}
#[derive(Debug, Deserialize, Serialize, ToSchema)]
pub struct ApiMessage {
pub role: ApiRole,
pub content: String,
}
#[derive(Debug, Deserialize, Serialize, ToSchema, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum ApiRole {
System,
User,
Assistant,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct ApiChatResponse {
pub id: Uuid,
pub object: ApiCompletionObject,
pub created: String,
pub model: String,
pub choices: Vec<ApiChatChoice>,
pub usage: Usage,
pub conversation_id: Uuid,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct ApiChatChoice {
pub index: u32,
pub message: ApiMessage,
pub finish_reason: ApiFinishReason,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct ChatCompletionChunk {
pub id: Uuid,
pub object: String,
pub choices: Vec<ChatChunkChoice>,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct ChatChunkChoice {
pub index: u32,
pub delta: Delta,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct Delta {
pub content: Option<String>,
pub role: Option<ApiRole>,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct StartEventData {
pub conversation_id: Uuid,
pub created: String,
pub id: Uuid,
pub model: String,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
pub struct EndEventData {
pub id: Uuid,
pub usage: Usage,
pub finish_reason: ApiFinishReason,
}
#[derive(Debug, Serialize, Deserialize, ToSchema)]
#[serde(tag = "type", content = "data")]
pub enum StreamEvent {
#[serde(rename = "start")]
Start(StartEventData),
#[serde(rename = "end")]
End(EndEventData),
#[serde(rename = "delta")]
Delta(ChatCompletionChunk),
}
// ------ Api Key ------
#[derive(Debug, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ApiKeyScope {
User,
Admin,
}
#[derive(serde::Deserialize)]
pub struct CreateApiKeyRequest {
pub name: String,
pub scopes: Vec<ApiKeyScope>,
}
#[derive(serde::Serialize)]
pub struct CreateApiKeyResponse {
pub api_key: String,
}
// ------ Fetch database ------
// --- Shared ---
#[derive(Debug, Deserialize)]
pub struct CursorPage {
pub limit: Option<u32>,
pub before: Option<chrono::DateTime<chrono::Utc>>,
}
// --- Conversation ---
#[derive(Debug, Serialize)]
pub struct ConversationSummary {
pub id: Uuid,
pub title: String,
pub created_at: chrono::DateTime<chrono::Utc>,
pub updated_at: chrono::DateTime<chrono::Utc>,
}
#[derive(Debug, Serialize)]
pub struct ConversationListResponse {
pub conversations: Vec<ConversationSummary>,
pub has_more: bool,
}
// --- Message ---
#[derive(Debug, Serialize)]
pub enum ApiChatRole {
System,
User,
Assistant,
}
#[derive(Debug, Serialize)]
pub struct MessageSummary {
pub id: Uuid,
pub parent_id: Option<Uuid>,
pub role: ApiChatRole,
pub content: String,
pub created_at: chrono::DateTime<chrono::Utc>,
pub tokens: Option<i32>,
}
#[derive(Debug, Serialize)]
pub struct MessageListResponse {
pub messages: Vec<MessageSummary>,
pub has_more: bool,
}
-59
View File
@@ -1,59 +0,0 @@
use jsonwebtoken::{Algorithm, DecodingKey, Validation, decode, decode_header};
use once_cell::sync::Lazy;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::env;
#[derive(Debug, Deserialize, Serialize, Clone)]
pub struct Claims {
pub sub: String,
pub preferred_username: Option<String>,
pub exp: usize,
pub iss: String,
pub aud: Option<String>,
pub realm_access: Option<RealmAccess>,
}
#[derive(Debug, Deserialize, Serialize, Clone)]
pub struct RealmAccess {
pub roles: Vec<String>,
}
static ISSUER: Lazy<String> = Lazy::new(|| env::var("ISSUER").expect("ISSUER not set"));
pub fn validate_token(token: &str, jwks: &Value) -> Result<Claims, String> {
// 1. Decode header
let header = decode_header(token).map_err(|_| "Invalid header")?;
let kid = header.kid.ok_or("Missing kid")?;
// 2. Find matching key
let keys = jwks["keys"].as_array().ok_or("Invalid JWKS")?;
let key = keys
.iter()
.find(|k| k["kid"] == kid)
.ok_or("Matching key not found")?;
// 3. Extract RSA components
let n = key["n"].as_str().ok_or("Missing n")?;
let e = key["e"].as_str().ok_or("Missing e")?;
let decoding_key =
DecodingKey::from_rsa_components(n, e).map_err(|_| "Invalid decoding key")?;
// 4. Setup validation rules
let mut validation = Validation::new(Algorithm::RS256);
validation.set_issuer(&[ISSUER.as_str()]);
// Optional but recommended:
validation.validate_exp = true;
validation.validate_aud = false; // depends on your Keycloak config
// 5. Decode & verify
let token_data = decode::<Claims>(token, &decoding_key, &validation)
.map_err(|_| "Token validation failed")?;
Ok(token_data.claims)
}
-47
View File
@@ -1,47 +0,0 @@
use axum::{extract::Request, http::StatusCode, middleware::Next, response::Response};
use crate::auth::{
jwt::validate_token,
keycloak::{get_jwks, refresh_jwks},
};
pub async fn auth_middleware(mut request: Request, next: Next) -> Result<Response, StatusCode> {
dbg!("Middleware hit");
let headers = request.headers();
dbg!("Headers extracted");
let auth_header = headers.get("authorization").and_then(|v| v.to_str().ok());
dbg!("Auth header: {:?}", auth_header);
let auth_header = auth_header.ok_or(StatusCode::UNAUTHORIZED)?;
let token = auth_header
.strip_prefix("Bearer ")
.ok_or(StatusCode::UNAUTHORIZED)?;
let jwks = get_jwks().await.map_err(|_| StatusCode::UNAUTHORIZED)?;
match validate_token(token, &jwks) {
Ok(claims) => {
dbg!("Token valid");
request.extensions_mut().insert(claims);
Ok(next.run(request).await)
}
Err(_) => {
let jwks = refresh_jwks().await.map_err(|_| StatusCode::UNAUTHORIZED)?;
match validate_token(token, &jwks) {
Ok(claims) => {
request.extensions_mut().insert(claims);
Ok(next.run(request).await)
}
Err(_) => Err(StatusCode::UNAUTHORIZED),
}
}
}
}
-3
View File
@@ -1,3 +0,0 @@
pub mod jwt;
pub mod keycloak;
pub mod middleware;
+27
View File
@@ -0,0 +1,27 @@
use uuid::Uuid;
#[derive(Debug)]
pub struct CreateApiKeyRequest {
pub user_id: Uuid,
pub name: String,
pub roles: Vec<KeyRole>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum KeyRole {
User,
Admin,
}
#[derive(Debug, Clone)]
pub struct AuthContext {
pub user_id: Uuid,
pub _api_key_id: Uuid,
pub roles: Vec<KeyRole>,
}
impl AuthContext {
pub fn has_role(&self, role: &KeyRole) -> bool {
self.roles.iter().any(|r| r == role)
}
}
+25
View File
@@ -0,0 +1,25 @@
// src/core/auth/jwt.rs
use std::collections::HashMap;
use uuid::Uuid;
#[derive(Debug, Clone)]
pub struct JwtClaims {
pub user_id: Uuid,
pub _username: Option<String>,
pub _exp: usize,
pub _issuer: String,
pub realm_roles: Vec<String>,
pub client_roles: HashMap<String, Vec<String>>,
}
impl JwtClaims {
pub fn has_role(&self, role: &str) -> bool {
self.realm_roles.iter().any(|r| r == role)
|| self
.client_roles
.get("chat-api")
.is_some_and(|roles| roles.iter().any(|r| r == role))
}
}
+33
View File
@@ -0,0 +1,33 @@
use uuid::Uuid;
pub mod api_key;
pub mod jwt;
#[derive(Debug, Clone)]
pub enum Auth {
Jwt(jwt::JwtClaims),
ApiKey(api_key::AuthContext),
}
impl Auth {
pub fn user_id(&self) -> Uuid {
match self {
Auth::Jwt(c) => c.user_id,
Auth::ApiKey(c) => c.user_id,
}
}
pub fn has_jwt_role(&self, role: &str) -> bool {
match self {
Auth::Jwt(c) => c.has_role(role),
Auth::ApiKey(_) => false,
}
}
pub fn has_apikey_role(&self, role: &crate::core::auth::api_key::KeyRole) -> bool {
match self {
Auth::Jwt(_) => false,
Auth::ApiKey(c) => c.has_role(role),
}
}
}
+44
View File
@@ -0,0 +1,44 @@
use crate::core::llm::ChatRole;
use uuid::Uuid;
#[derive(Debug)]
pub struct ConversationSummary {
pub id: Uuid,
pub title: String,
pub created_at: chrono::DateTime<chrono::Utc>,
pub updated_at: chrono::DateTime<chrono::Utc>,
}
#[derive(Debug)]
pub struct ConversationList {
pub conversations: Vec<ConversationSummary>,
pub has_more: bool,
}
pub struct CursorPage {
pub limit: u32,
pub before: Option<chrono::DateTime<chrono::Utc>>,
}
#[derive(Debug)]
pub struct MessageSummary {
pub id: Uuid,
pub parent_id: Option<Uuid>,
pub role: ChatRole,
pub content: String,
pub created_at: chrono::DateTime<chrono::Utc>,
pub tokens: Option<i32>,
}
#[derive(Debug)]
pub struct MessageList {
pub messages: Vec<MessageSummary>,
pub has_more: bool,
}
#[derive(Debug)]
pub enum ConversationResult {
Existing(Uuid),
Created(Uuid),
}
+1
View File
@@ -0,0 +1 @@
pub mod conversations;
+98
View File
@@ -0,0 +1,98 @@
use super::ChatRole;
use crate::services::errors::ServiceError;
use futures::Stream;
use serde::Serialize;
use std::pin::Pin;
use uuid::Uuid;
#[derive(Debug, Clone, Default)]
pub struct ChatCompletionOptions {
pub seed: Option<i64>,
pub temperature: Option<f32>,
pub top_p: Option<f32>,
pub top_k: Option<u32>,
pub stop: Option<Vec<String>>,
pub num_ctx: Option<u32>,
pub num_predict: Option<u32>,
pub keep_alive: Option<String>,
pub stream: bool,
pub context_depth: u32,
}
#[derive(Debug, Clone, Serialize)]
pub struct Message {
pub role: ChatRole,
pub content: String,
}
#[derive(Debug, Clone)]
pub struct ChatCompletionRequest {
pub model: String,
pub message: Message,
pub options: ChatCompletionOptions,
pub conversation_id: Option<uuid::Uuid>,
pub parent_id: Option<uuid::Uuid>,
}
#[derive(Debug, Clone, Serialize)]
pub struct ChatCompletionResultNoStream {
pub model: String,
pub message: Message,
pub created_at: String,
pub id: uuid::Uuid,
pub conversation_id: Uuid,
pub prompt_tokens: u32,
pub completion_tokens: u32,
pub finish_reason: FinishReason,
pub total_duration: u64,
pub load_duration: u64,
}
#[derive(Debug, Serialize, Clone)]
pub enum FinishReason {
Stop,
Length,
Error,
}
#[derive(Debug, Serialize)]
pub enum ChatCompletionStreamEvent {
Start {
conversation_id: Uuid,
message_id: Uuid,
created_at: String,
model: String,
},
Token {
content: String,
id: uuid::Uuid,
created_at: String,
},
Final(ChatCompletionResultNoStream),
}
pub type ChatCompletionStream =
Pin<Box<dyn Stream<Item = Result<ChatCompletionStreamEvent, ServiceError>> + Send>>;
pub enum ChatCompletionResult {
Stream(ChatCompletionStream),
NoStream(ChatCompletionResultNoStream),
}
impl Message {
pub fn from_summary(summary: crate::core::databases::conversations::MessageSummary) -> Self {
Self {
role: summary.role,
content: summary.content,
}
}
}
+58
View File
@@ -0,0 +1,58 @@
use crate::providers::ollama::errors::LlmError;
use futures::Stream;
use serde::Serialize;
use std::pin::Pin;
#[derive(Debug, Clone, Default)]
pub struct CompletionOptions {
pub seed: Option<i64>,
pub temperature: Option<f32>,
pub top_p: Option<f32>,
pub top_k: Option<u32>,
pub stop: Option<Vec<String>>,
pub num_ctx: Option<u32>,
pub num_predict: Option<u32>,
pub keep_alive: Option<String>,
pub stream: bool,
}
#[derive(Debug, Clone)]
pub struct CompletionRequest {
pub model: String,
pub prompt: String,
pub options: CompletionOptions,
}
#[derive(Debug, Clone, Serialize)]
pub struct CompletionResultNoStream {
pub model: String,
pub text: String,
pub created_at: String,
pub id: uuid::Uuid,
pub prompt_tokens: u32,
pub completion_tokens: u32,
pub done_reason: Option<String>,
pub total_duration: Option<u64>,
pub load_duration: Option<u64>,
}
#[derive(Serialize)]
pub enum CompletionStreamEvent {
Token(String),
Final(CompletionResultNoStream),
}
pub type CompletionStream =
Pin<Box<dyn Stream<Item = Result<CompletionStreamEvent, LlmError>> + Send>>;
pub enum CompletionResult {
Stream(CompletionStream),
NoStream(CompletionResultNoStream),
}
+6
View File
@@ -0,0 +1,6 @@
pub mod chat;
pub mod completions;
pub mod models;
pub mod role;
pub use role::ChatRole;
+46
View File
@@ -0,0 +1,46 @@
use std::collections::HashMap;
#[derive(Debug, Clone)]
pub struct Models {
pub models: Vec<Model>,
}
#[derive(Debug, Clone)]
pub struct Model {
pub name: String,
/// Optional grouping (OpenAI = "gpt", Ollama = "llama", etc.)
pub family: Option<String>,
/// Human-readable size like "7B", "13B", "gpt-4"
pub size: Option<String>,
/// Optional metadata, provider-specific info normalized into a string map
pub metadata: ModelMetadata,
}
#[derive(Debug, Clone, Default)]
pub struct ModelMetadata {
pub extra: HashMap<String, String>,
}
#[derive(Debug, Clone)]
pub struct LoadModelRequest {
pub model: String,
pub keep_alive: String,
}
#[derive(Debug, Clone)]
pub struct LoadModelResponse {
pub model: String,
}
#[derive(Debug, Clone)]
pub struct UnloadModelRequest {
pub model: String,
}
#[derive(Debug, Clone)]
pub struct UnloadModelResponse {
pub model: String,
}
+8
View File
@@ -0,0 +1,8 @@
use serde::Serialize;
#[derive(Debug, Clone, Serialize)]
pub enum ChatRole {
System,
User,
Assistant,
}
+3
View File
@@ -0,0 +1,3 @@
pub mod auth;
pub mod databases;
pub mod llm;
+16
View File
@@ -0,0 +1,16 @@
use thiserror::Error;
#[derive(Debug, Error)]
pub enum DbError {
#[error("database connection error")]
Connection(#[from] sqlx::Error),
#[error("database timeout")]
Timeout,
#[error("not authorized")]
Unauthorized,
#[error("not found")]
NotFound,
}
+3
View File
@@ -0,0 +1,3 @@
pub mod postgres;
pub mod errors;
+2
View File
@@ -0,0 +1,2 @@
pub mod queries;
pub mod types;
+59
View File
@@ -0,0 +1,59 @@
use super::types;
use super::types::Role;
use crate::databases::errors::DbError;
use sqlx::PgPool;
use uuid::Uuid;
pub async fn update_last_access(pool: &PgPool, api_key_id: &Uuid) -> Result<(), DbError> {
sqlx::query!(
r#"
UPDATE auth.api_key
SET last_used_at = now()
WHERE id = $1
"#,
api_key_id
)
.execute(pool)
.await?;
Ok(())
}
pub async fn create(pool: &PgPool, key: super::types::CreateApiKey) -> Result<(), DbError> {
sqlx::query!(
r#"
INSERT INTO auth.api_key (key_hash, name, created_by, scopes)
VALUES ($1, $2, $3, $4::auth.role[])
"#,
key.key_hash,
key.name,
key.user_id,
&key.roles as _
)
.execute(pool)
.await?;
Ok(())
}
pub async fn validate(pool: &PgPool, key_hash: &str) -> Result<types::AuthContext, DbError> {
let row = sqlx::query_as!(
types::AuthContext,
r#"
SELECT
u.id AS user_id,
ak.id AS api_key_id,
ak.scopes AS "roles!: Vec<Role>"
FROM auth.api_key ak
JOIN auth.app_user u ON u.id = ak.created_by
WHERE ak.key_hash = $1
AND ak.revoked_at IS NULL
"#,
key_hash
)
.fetch_optional(pool)
.await?;
row.ok_or(DbError::Unauthorized)
}
+23
View File
@@ -0,0 +1,23 @@
use sqlx::Type;
#[derive(Debug)]
pub struct CreateApiKey {
pub key_hash: String,
pub name: String,
pub user_id: uuid::Uuid,
pub roles: Vec<Role>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Type)]
#[sqlx(type_name = "auth.role", rename_all = "lowercase")]
pub enum Role {
User,
Admin,
}
#[derive(Debug, sqlx::FromRow)]
pub struct AuthContext {
pub user_id: uuid::Uuid,
pub api_key_id: uuid::Uuid,
pub roles: Vec<Role>,
}
+2
View File
@@ -0,0 +1,2 @@
pub mod queries;
pub mod types;
+231
View File
@@ -0,0 +1,231 @@
use crate::databases::errors::DbError;
use crate::databases::postgres::chat::{types, types::MessageRole};
use sqlx::{Acquire, PgPool};
use uuid::Uuid;
// ---- Helpers ----
async fn validate_conversation<'e, E>(executor: E, conversation_id: Uuid) -> Result<(), DbError>
where
E: sqlx::Executor<'e, Database = sqlx::Postgres>,
{
tracing::debug!("Validating Conversation");
let rec = sqlx::query!(
r#"SELECT id FROM chat.conversation WHERE id = $1"#,
conversation_id
)
.fetch_optional(executor)
.await?;
match rec {
Some(_) => Ok(()),
None => Err(DbError::NotFound),
}
}
async fn create_conversation<'e, E>(executor: E, user_id: Uuid) -> Result<Uuid, DbError>
where
E: sqlx::Executor<'e, Database = sqlx::Postgres>,
{
tracing::debug!("Creating Conversation");
let rec = sqlx::query!(
r#"INSERT INTO chat.conversation (user_id) VALUES ($1) RETURNING id"#,
user_id
)
.fetch_one(executor)
.await?;
Ok(rec.id)
}
// ------ Creation ------
pub async fn get_or_create_conversation(
pool: &PgPool,
conversation_id: Option<Uuid>,
user_id: Uuid,
) -> Result<types::ConversationState, DbError> {
let mut tx: sqlx::Transaction<'_, sqlx::Postgres> = pool.begin().await?;
let result = {
let conn = tx.acquire().await?;
sqlx::query(r#"SELECT set_config('app.current_user_id', $1::text, true)"#)
.bind(user_id.to_string())
.execute(&mut *conn)
.await?;
match conversation_id {
Some(id) => {
validate_conversation(&mut *conn, id).await?;
types::ConversationState::Existing(id)
}
None => {
types::ConversationState::Created(create_conversation(&mut *conn, user_id).await?)
}
}
};
tx.commit().await?;
Ok(result)
}
pub async fn set_conversation_title(
pool: &PgPool,
conversation_id: Uuid,
user_id: Uuid,
title: &str,
) -> Result<(), DbError> {
let mut tx = pool.begin().await?;
let conn = tx.acquire().await?;
sqlx::query(r#"SELECT set_config('app.current_user_id', $1::text, true)"#)
.bind(user_id.to_string())
.execute(&mut *conn)
.await?;
sqlx::query!(
r#"UPDATE chat.conversation SET title = $1 WHERE id = $2"#,
title,
conversation_id
)
.execute(&mut *conn)
.await?;
tx.commit().await?;
Ok(())
}
pub async fn insert_message(
pool: &PgPool,
user_id: Uuid,
conversation_id: Uuid,
parent_id: Option<Uuid>,
role: types::MessageRole,
content: &str,
tokens: Option<u32>,
) -> Result<Uuid, DbError> {
let mut tx = pool.begin().await?;
let conn = tx.acquire().await?;
sqlx::query(r#"SELECT set_config('app.current_user_id', $1::text, true)"#)
.bind(user_id.to_string())
.execute(&mut *conn)
.await?;
let row = sqlx::query!(
r#"
INSERT INTO chat.message (conversation_id, parent_id, role, content, tokens)
VALUES ($1, $2, $3, $4, $5)
RETURNING id
"#,
conversation_id,
parent_id,
role as types::MessageRole,
content,
tokens.unwrap_or(0) as i32
)
.fetch_one(&mut *conn)
.await?;
tx.commit().await?;
Ok(row.id)
}
pub async fn update_message_tokens(
pool: &PgPool,
user_id: Uuid,
message_id: Uuid,
tokens: u32,
) -> Result<(), DbError> {
let mut tx = pool.begin().await?;
let conn = tx.acquire().await?;
sqlx::query(r#"SELECT set_config('app.current_user_id', $1::text, true)"#)
.bind(user_id.to_string())
.execute(&mut *conn)
.await?;
sqlx::query!(
r#"UPDATE chat.message SET tokens = $1 WHERE id = $2"#,
tokens as i32,
message_id
)
.execute(&mut *conn)
.await?;
tx.commit().await?;
Ok(())
}
// ---- Fetching ----
pub async fn get_conversations_entries(
pool: &PgPool,
user_id: Uuid,
limit: u32,
before: Option<chrono::DateTime<chrono::Utc>>,
) -> Result<Vec<types::ConversationSummary>, DbError> {
let mut tx = pool.begin().await?;
sqlx::query("SELECT set_config('app.current_user_id', $1::text, true)")
.bind(user_id.to_string())
.execute(&mut *tx)
.await?;
let rows = sqlx::query_as!(
types::ConversationSummary,
r#"
SELECT id, title, created_at, updated_at
FROM chat.conversation
WHERE user_id = $1
AND ($2::timestamptz IS NULL OR updated_at < $2)
ORDER BY updated_at DESC
LIMIT $3
"#,
user_id,
before,
limit as i64
)
.fetch_all(&mut *tx)
.await?;
tx.commit().await?;
Ok(rows)
}
pub async fn get_conversation_messages(
pool: &PgPool,
user_id: Uuid,
conversation_id: Uuid,
limit: u32,
before: Option<chrono::DateTime<chrono::Utc>>,
) -> Result<Vec<types::MessageSummary>, DbError> {
let mut conn = pool.acquire().await?;
sqlx::query(r#"SELECT set_config('app.current_user_id', $1::text, true)"#)
.bind(user_id.to_string())
.execute(&mut *conn)
.await?;
let rows = sqlx::query_as!(
types::MessageSummary,
r#"
SELECT id, parent_id, role as "role: MessageRole", content, created_at, tokens
FROM chat.message
WHERE conversation_id = $1
AND ($2::timestamptz IS NULL OR created_at < $2)
ORDER BY created_at DESC
LIMIT $3
"#,
conversation_id,
before,
limit as i64
)
.fetch_all(&mut *conn)
.await?;
Ok(rows)
}
+34
View File
@@ -0,0 +1,34 @@
use uuid::Uuid;
#[derive(Debug, Clone, sqlx::Type)]
#[sqlx(type_name = "chat.role")]
#[sqlx(rename_all = "lowercase")]
pub enum MessageRole {
User,
Assistant,
System,
}
#[derive(Debug)]
pub enum ConversationState {
Existing(Uuid),
Created(Uuid),
}
#[derive(Debug, sqlx::FromRow)]
pub struct ConversationSummary {
pub id: Uuid,
pub title: Option<String>,
pub created_at: chrono::DateTime<chrono::Utc>,
pub updated_at: chrono::DateTime<chrono::Utc>,
}
#[derive(Debug, sqlx::FromRow)]
pub struct MessageSummary {
pub id: Uuid,
pub parent_id: Option<Uuid>,
pub role: MessageRole,
pub content: String,
pub created_at: chrono::DateTime<chrono::Utc>,
pub tokens: Option<i32>,
}
+5
View File
@@ -0,0 +1,5 @@
pub mod pool;
pub mod api_key;
pub mod chat;
pub mod user_activity;
+19
View File
@@ -0,0 +1,19 @@
use crate::databases::errors::DbError;
use sqlx::{PgPool, postgres::PgPoolOptions};
use std::time::Duration;
pub async fn create_pool(database_url: &str) -> Result<PgPool, DbError> {
PgPoolOptions::new()
.max_connections(10)
.acquire_timeout(Duration::from_secs(5))
.connect(database_url)
.await
.map_err(|err| {
tracing::error!("Postgres connection error: {:?}", err);
match err {
sqlx::Error::PoolTimedOut => DbError::Timeout,
_ => DbError::Connection(err),
}
})
}
@@ -0,0 +1 @@
pub mod queries;
@@ -0,0 +1,21 @@
use crate::databases::errors::DbError;
use sqlx::PgPool;
use uuid::Uuid;
pub async fn upsert_user_activity(pool: &PgPool, user_id: &Uuid) -> Result<(), DbError> {
sqlx::query!(
r#"
INSERT INTO auth.app_user (id, last_seen_at)
VALUES ($1, now())
ON CONFLICT (id)
DO UPDATE SET last_seen_at = now()
"#,
user_id
)
.execute(pool)
.await
.map_err(DbError::Connection)?;
Ok(())
}
+6
View File
@@ -0,0 +1,6 @@
pub mod api;
pub mod core;
pub mod databases;
pub mod mappers;
pub mod providers;
pub mod services;
+12 -49
View File
@@ -1,61 +1,24 @@
mod auth; mod api;
mod core;
mod databases;
mod mappers;
mod providers;
mod services;
use crate::auth::jwt::Claims; use api::app::build_app;
use crate::auth::middleware::auth_middleware;
use axum::extract::Extension;
use axum::{Router, middleware, routing::get};
use std::net::SocketAddr; use std::net::SocketAddr;
pub async fn protected_route(Extension(claims): Extension<Claims>) -> String {
format!(
"Hello {}, your user id is {}",
claims.preferred_username.unwrap_or("unknown".to_string()),
claims.sub
)
}
pub async fn public_route() -> &'static str {
println!("Public route hit");
"Public endpoint: no authentication required"
}
pub fn app() -> Router {
let public_routes = Router::new().route("/", get(public_route));
let protected_routes = Router::new()
.route("/protected", get(protected_route))
.layer(middleware::from_fn(auth_middleware));
Router::new().merge(public_routes).merge(protected_routes)
}
#[tokio::main] #[tokio::main]
async fn main() { async fn main() {
#[cfg(debug_assertions)] #[cfg(debug_assertions)]
{ dotenvy::dotenv().ok();
dotenvy::dotenv().ok();
}
// let jwks = keycloak::get_jwks() tracing_subscriber::fmt::init();
// .await
// .expect("Failed to fetch JWKS");
// // println!("{:?}", jwks); let app = build_app().await;
// let token = "eyJhbGciOiJSUzI1NiIsInR5cCIgOiAiSldUIiwia2lkIiA6ICJublpLek04TkZHVmpWbGFPRXZpMUtFSTVHQWRwaGlsYjh3RHRLeG5JOENZIn0.eyJleHAiOjE3NzU3Mzg2MDUsImlhdCI6MTc3NTczODMwNSwianRpIjoiNjk2OTY4NzQtZWMwNi00NGFkLTg0MDYtYmY3YWM4MjI5MjkxIiwiaXNzIjoiaHR0cHM6Ly9hdXRoLmljZWJlcmcuYmxhY2svcmVhbG1zL2ljZWJlcmciLCJzdWIiOiJmZGRiN2FjZC1kMmE5LTRmMTctOWIxNi1kZjVlN2EzNDI4YjciLCJ0eXAiOiJCZWFyZXIiLCJhenAiOiJjaGF0LWFwaSIsInNjb3BlIjoiIiwiY2xpZW50SG9zdCI6Ijg2LjIxMi44NC4xOTEiLCJjbGllbnRBZGRyZXNzIjoiODYuMjEyLjg0LjE5MSIsImNsaWVudF9pZCI6ImNoYXQtYXBpIn0.BS7ohLWiMDxAUz_Q-Qi2UoLYbNn8AUrYeSWeO-602SQ-AYBW3gfYxXOSeRgWyn4VfObpVfK7QfqQBUxorXxi1JVld-4fGXL8NXQNyq5Ip_JHNG1p02Z39Pe9MmC9MXOwA_GQF2PIkLIdOJ_W_guXVhl2ptEWPPSiXM5Z5CNg8lyOiKPI0g2JWV6FBRG-HMXzqnxAb1j8wGUpC9JzGwAU3sjWBGhT1AAovs-XLmm5hZEPxI-Ia3SmUnF-QjFMmebPVxLdxL7OszzVEhKipsZRiwQxjY6eJhJFFa8uycBigHPSzu_HqqkK6AjNlyExvR0EGvl9zUWdOfMPDiVX2Sg92g"; let addr = SocketAddr::from(([0, 0, 0, 0], 3001));
// match jwt::validate_token(token, &jwks) { tracing::debug!("Server running on {}", addr);
// Ok(claims) => {
// println!("Valid token for user: {:?}", claims);
// }
// Err(err) => {
// println!("Invalid token: {}", err);
// }
// }
let app = app();
let addr = SocketAddr::from(([0, 0, 0, 0], 3000));
println!("Server running on {}", addr);
axum::serve(tokio::net::TcpListener::bind(addr).await.unwrap(), app) axum::serve(tokio::net::TcpListener::bind(addr).await.unwrap(), app)
.await .await
+66
View File
@@ -0,0 +1,66 @@
use crate::{api, core};
impl From<api::types::ApiCompletionRequest> for core::llm::completions::CompletionRequest {
fn from(m: api::types::ApiCompletionRequest) -> Self {
Self {
model: m.model,
prompt: m.prompt,
options: core::llm::completions::CompletionOptions {
temperature: m.options.temperature,
keep_alive: m.options.keep_alive,
num_ctx: m.options.num_ctx,
top_p: m.options.top_p,
top_k: m.options.top_k,
seed: m.options.seed,
num_predict: m.options.num_predict,
stop: m.options.stop,
stream: m.options.stream,
},
}
}
}
impl From<api::types::ApiChatRequest> for core::llm::chat::ChatCompletionRequest {
fn from(m: api::types::ApiChatRequest) -> Self {
Self {
model: m.model,
message: m.message.into(),
options: core::llm::chat::ChatCompletionOptions {
temperature: m.base.temperature,
keep_alive: m.base.keep_alive,
num_ctx: m.base.num_ctx,
top_p: m.base.top_p,
top_k: m.base.top_k,
seed: m.base.seed,
num_predict: m.base.num_predict,
stop: m.base.stop,
stream: m.base.stream,
context_depth: m.base.context_depth.unwrap_or(10),
},
conversation_id: m.conversation_id,
parent_id: m.parent_id,
}
}
}
impl From<api::types::ApiMessage> for core::llm::chat::Message {
fn from(m: api::types::ApiMessage) -> Self {
Self {
role: match m.role {
api::types::ApiRole::System => core::llm::ChatRole::System,
api::types::ApiRole::User => core::llm::ChatRole::User,
api::types::ApiRole::Assistant => core::llm::ChatRole::Assistant,
},
content: m.content,
}
}
}
impl From<api::types::ApiKeyScope> for core::auth::api_key::KeyRole {
fn from(role: api::types::ApiKeyScope) -> Self {
match role {
api::types::ApiKeyScope::User => core::auth::api_key::KeyRole::User,
api::types::ApiKeyScope::Admin => core::auth::api_key::KeyRole::Admin,
}
}
}
+146
View File
@@ -0,0 +1,146 @@
use crate::{api, core};
impl From<core::llm::models::Models> for api::types::ApiModelsResponse {
fn from(m: core::llm::models::Models) -> Self {
Self {
models: m.models.into_iter().map(Into::into).collect(),
}
}
}
impl From<core::llm::models::ModelMetadata> for api::types::ApiModelMetadata {
fn from(m: core::llm::models::ModelMetadata) -> Self {
Self {
extra: m
.extra
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect(),
}
}
}
impl From<core::llm::models::Model> for api::types::ApiModelInfo {
fn from(m: core::llm::models::Model) -> Self {
Self {
name: m.name,
family: m.family,
parameter_size: m.size,
metadata: m.metadata.into(),
}
}
}
impl From<core::databases::conversations::ConversationSummary> for api::types::ConversationSummary {
fn from(c: core::databases::conversations::ConversationSummary) -> Self {
Self {
id: c.id,
title: c.title,
created_at: c.created_at,
updated_at: c.updated_at,
}
}
}
impl From<core::databases::conversations::ConversationList>
for api::types::ConversationListResponse
{
fn from(c: core::databases::conversations::ConversationList) -> Self {
Self {
conversations: c.conversations.into_iter().map(Into::into).collect(),
has_more: c.has_more,
}
}
}
impl From<core::databases::conversations::MessageSummary> for api::types::MessageSummary {
fn from(c: core::databases::conversations::MessageSummary) -> Self {
Self {
id: c.id,
parent_id: c.parent_id,
role: match c.role {
core::llm::ChatRole::User => api::types::ApiChatRole::User,
core::llm::ChatRole::Assistant => api::types::ApiChatRole::Assistant,
core::llm::ChatRole::System => api::types::ApiChatRole::System,
},
content: c.content,
created_at: c.created_at,
tokens: c.tokens,
}
}
}
impl From<core::databases::conversations::MessageList> for api::types::MessageListResponse {
fn from(c: core::databases::conversations::MessageList) -> Self {
Self {
messages: c.messages.into_iter().map(Into::into).collect(),
has_more: c.has_more,
}
}
}
impl From<core::llm::chat::Message> for api::types::ApiMessage {
fn from(c: core::llm::chat::Message) -> Self {
Self {
role: match c.role {
core::llm::ChatRole::User => api::types::ApiRole::User,
core::llm::ChatRole::Assistant => api::types::ApiRole::Assistant,
core::llm::ChatRole::System => api::types::ApiRole::System,
},
content: c.content,
}
}
}
impl From<core::llm::completions::CompletionResultNoStream> for api::types::ApiCompletionResponse {
fn from(m: core::llm::completions::CompletionResultNoStream) -> Self {
Self {
id: m.id,
object: api::types::ApiCompletionObject::TextCompletion,
created: m.created_at,
model: m.model,
choices: vec![api::types::Choice {
finish_reason: api::types::ApiFinishReason::Stop,
text: m.text,
index: 0,
}],
usage: api::types::Usage {
completion_tokens: m.completion_tokens,
prompt_tokens: m.prompt_tokens,
total_tokens: m.completion_tokens + m.prompt_tokens,
},
}
}
}
impl From<core::llm::chat::ChatCompletionResultNoStream> for api::types::ApiChatResponse {
fn from(m: core::llm::chat::ChatCompletionResultNoStream) -> Self {
Self {
id: m.id,
conversation_id: m.conversation_id,
object: api::types::ApiCompletionObject::TextCompletion,
created: m.created_at,
model: m.model,
choices: vec![api::types::ApiChatChoice {
finish_reason: api::types::ApiFinishReason::Stop,
message: m.message.into(),
index: 0,
}],
usage: api::types::Usage {
completion_tokens: m.completion_tokens,
prompt_tokens: m.prompt_tokens,
total_tokens: m.completion_tokens + m.prompt_tokens,
},
}
}
}
impl From<core::llm::chat::FinishReason> for api::types::ApiFinishReason {
fn from(f: core::llm::chat::FinishReason) -> Self {
match f {
core::llm::chat::FinishReason::Stop => Self::Stop,
core::llm::chat::FinishReason::Length => Self::Length,
core::llm::chat::FinishReason::Error => Self::Error,
}
}
}
+30
View File
@@ -0,0 +1,30 @@
impl From<crate::core::llm::ChatRole> for crate::databases::postgres::chat::types::MessageRole {
fn from(m: crate::core::llm::ChatRole) -> Self {
match m {
crate::core::llm::ChatRole::System => {
crate::databases::postgres::chat::types::MessageRole::System
}
crate::core::llm::ChatRole::Assistant => {
crate::databases::postgres::chat::types::MessageRole::Assistant
}
crate::core::llm::ChatRole::User => {
crate::databases::postgres::chat::types::MessageRole::User
}
}
}
}
impl From<crate::core::auth::api_key::KeyRole>
for crate::databases::postgres::api_key::types::Role
{
fn from(role: crate::core::auth::api_key::KeyRole) -> Self {
match role {
crate::core::auth::api_key::KeyRole::User => {
crate::databases::postgres::api_key::types::Role::User
}
crate::core::auth::api_key::KeyRole::Admin => {
crate::databases::postgres::api_key::types::Role::Admin
}
}
}
}
+67
View File
@@ -0,0 +1,67 @@
use crate::core;
use crate::providers::ollama;
impl From<core::llm::completions::CompletionOptions> for ollama::types::OllamaOptions {
fn from(opts: core::llm::completions::CompletionOptions) -> Self {
Self {
seed: opts.seed,
temperature: opts.temperature,
top_p: opts.top_p,
top_k: opts.top_k,
stop: opts.stop,
num_ctx: opts.num_ctx,
num_predict: opts.num_predict,
}
}
}
impl From<core::llm::chat::ChatCompletionOptions> for ollama::types::OllamaOptions {
fn from(opts: core::llm::chat::ChatCompletionOptions) -> Self {
Self {
seed: opts.seed,
temperature: opts.temperature,
top_p: opts.top_p,
top_k: opts.top_k,
stop: opts.stop,
num_ctx: opts.num_ctx,
num_predict: opts.num_predict,
}
}
}
impl From<core::llm::chat::Message> for ollama::types::OllamaMessage {
fn from(m: core::llm::chat::Message) -> Self {
Self {
role: match m.role {
core::llm::ChatRole::System => ollama::types::OllamaRole::System,
core::llm::ChatRole::User => ollama::types::OllamaRole::User,
core::llm::ChatRole::Assistant => ollama::types::OllamaRole::Assistant,
},
content: m.content,
}
}
}
impl From<core::llm::completions::CompletionRequest> for ollama::types::OllamaGenerateRequest {
fn from(r: core::llm::completions::CompletionRequest) -> Self {
Self {
model: r.model,
prompt: r.prompt,
stream: r.options.stream,
keep_alive: r.options.keep_alive.clone().unwrap_or("5m".to_string()),
options: Some(r.options.clone().into()),
}
}
}
impl From<core::llm::chat::ChatCompletionRequest> for ollama::types::OllamaChatRequest {
fn from(r: core::llm::chat::ChatCompletionRequest) -> Self {
Self {
model: r.model,
messages: Vec::new(),
stream: r.options.stream,
keep_alive: r.options.keep_alive.clone().unwrap_or("5m".to_string()),
options: Some(r.options.clone().into()),
}
}
}
+69
View File
@@ -0,0 +1,69 @@
use crate::core;
use crate::databases::postgres;
impl From<postgres::chat::types::ConversationSummary>
for core::databases::conversations::ConversationSummary
{
fn from(c: postgres::chat::types::ConversationSummary) -> Self {
Self {
id: c.id,
title: c.title.unwrap_or("No title Generated".to_string()),
created_at: c.created_at,
updated_at: c.updated_at,
}
}
}
impl From<postgres::chat::types::MessageSummary>
for core::databases::conversations::MessageSummary
{
fn from(m: postgres::chat::types::MessageSummary) -> Self {
Self {
id: m.id,
parent_id: m.parent_id,
role: match m.role {
postgres::chat::types::MessageRole::User => core::llm::ChatRole::User,
postgres::chat::types::MessageRole::Assistant => core::llm::ChatRole::Assistant,
postgres::chat::types::MessageRole::System => core::llm::ChatRole::System,
},
content: m.content,
created_at: m.created_at,
tokens: m.tokens,
}
}
}
impl From<postgres::api_key::types::Role> for core::auth::api_key::KeyRole {
fn from(role: postgres::api_key::types::Role) -> Self {
match role {
postgres::api_key::types::Role::User => core::auth::api_key::KeyRole::User,
postgres::api_key::types::Role::Admin => core::auth::api_key::KeyRole::Admin,
}
}
}
impl From<postgres::api_key::types::AuthContext> for core::auth::api_key::AuthContext {
fn from(m: postgres::api_key::types::AuthContext) -> Self {
Self {
user_id: m.user_id,
_api_key_id: m.api_key_id,
roles: m.roles.into_iter().map(Into::into).collect(),
}
}
}
impl From<postgres::chat::types::ConversationState>
for core::databases::conversations::ConversationResult
{
fn from(m: postgres::chat::types::ConversationState) -> Self {
match m {
postgres::chat::types::ConversationState::Existing(id) => {
core::databases::conversations::ConversationResult::Existing(id)
}
postgres::chat::types::ConversationState::Created(id) => {
core::databases::conversations::ConversationResult::Created(id)
}
}
}
}
+30
View File
@@ -0,0 +1,30 @@
use crate::core::auth::jwt::JwtClaims;
use crate::providers::keycloak::claims::KeycloakClaims;
use std::collections::HashMap;
use uuid::Uuid;
impl From<KeycloakClaims> for JwtClaims {
fn from(c: KeycloakClaims) -> Self {
let realm_roles = c
.realm_access
.as_ref()
.map(|r| r.roles.clone())
.unwrap_or_default();
let client_roles = c
.resource_access
.into_iter()
.map(|(k, v)| (k, v.roles))
.collect::<HashMap<_, _>>();
Self {
user_id: c.sub.parse().unwrap_or_else(|_| Uuid::nil()),
_username: c.preferred_username,
_exp: c.exp,
_issuer: c.iss,
realm_roles,
client_roles,
}
}
}
+7
View File
@@ -0,0 +1,7 @@
pub mod api_to_core;
pub mod core_to_api;
pub mod core_to_database;
pub mod core_to_ollama;
pub mod database_to_core;
pub mod keycloak_to_core;
pub mod ollama_to_core;
+60
View File
@@ -0,0 +1,60 @@
use crate::core;
use crate::providers::ollama;
impl From<ollama::types::OllamaModels> for core::llm::models::Models {
fn from(m: ollama::types::OllamaModels) -> Self {
Self {
models: m.models.into_iter().map(Into::into).collect(),
}
}
}
impl From<ollama::types::OllamaModel> for core::llm::models::Model {
fn from(m: ollama::types::OllamaModel) -> Self {
Self {
name: m.name,
family: m.details.as_ref().and_then(|d| d.family.clone()),
size: m.details.as_ref().and_then(|d| d.parameter_size.clone()),
metadata: core::llm::models::ModelMetadata {
extra: [
("digest", m.digest.unwrap_or_default()),
("size_bytes", m.size.unwrap_or_default().to_string()),
("modified_at", m.modified_at.unwrap_or_default().to_string()),
(
"quantization_level",
m.details
.as_ref()
.and_then(|d| d.quantization_level.clone())
.unwrap_or_default(),
),
]
.into_iter()
.map(|(k, v)| (k.to_string(), v))
.collect(),
},
}
}
}
impl From<ollama::types::OllamaMessage> for core::llm::chat::Message {
fn from(m: ollama::types::OllamaMessage) -> Self {
Self {
role: match m.role {
ollama::types::OllamaRole::System => core::llm::ChatRole::System,
ollama::types::OllamaRole::User => core::llm::ChatRole::User,
ollama::types::OllamaRole::Assistant => core::llm::ChatRole::Assistant,
},
content: m.content,
}
}
}
impl From<ollama::types::OllamaFinishReason> for core::llm::chat::FinishReason {
fn from(f: ollama::types::OllamaFinishReason) -> Self {
match f {
ollama::types::OllamaFinishReason::Stop => Self::Stop,
ollama::types::OllamaFinishReason::Length => Self::Length,
ollama::types::OllamaFinishReason::Error => Self::Error,
}
}
}
+23
View File
@@ -0,0 +1,23 @@
use serde::Deserialize;
use std::collections::HashMap;
#[derive(Debug, Clone, Default, Deserialize)]
pub struct KeycloakClaims {
pub sub: String,
pub preferred_username: Option<String>,
pub exp: usize,
pub iss: String,
pub _aud: Option<Vec<String>>,
pub realm_access: Option<RealmAccess>,
pub resource_access: HashMap<String, ResourceAccess>,
}
#[derive(Debug, Clone, Deserialize)]
pub struct RealmAccess {
pub roles: Vec<String>,
}
#[derive(Debug, Clone, Deserialize)]
pub struct ResourceAccess {
pub roles: Vec<String>,
}
+40
View File
@@ -0,0 +1,40 @@
use thiserror::Error;
#[derive(Debug, Error)]
pub enum JwtValidationError {
#[error("invalid authorization header")]
InvalidHeader,
#[error("missing 'kid' in token header")]
MissingKid,
#[error("invalid JWKS structure")]
InvalidJwks,
#[error("no matching key found for kid")]
JwkNotFound,
#[error("missing RSA modulus")]
MissingModulus,
#[error("missing RSA exponent")]
MissingExponent,
#[error("invalid decoding key")]
InvalidDecodingKey,
#[error("token validation failed")]
TokenValidationFailed,
#[error("invalid or expired token")]
InvalidToken,
#[error("failed to fetch JWKS")]
JwksFetchFailed,
#[error("failed to refresh JWKS")]
JwksRefreshFailed,
#[error(transparent)]
Reqwest(#[from] reqwest::Error),
}
@@ -1,3 +1,5 @@
use super::errors::JwtValidationError;
use once_cell::sync::Lazy; use once_cell::sync::Lazy;
use serde_json::Value; use serde_json::Value;
use std::env; use std::env;
@@ -15,7 +17,7 @@ static JWK_CACHE: Lazy<Arc<RwLock<Option<JwksCache>>>> = Lazy::new(|| Arc::new(R
static JWKS_URL: Lazy<String> = Lazy::new(|| env::var("JWKS_URL").expect("JWKS_URL not set")); static JWKS_URL: Lazy<String> = Lazy::new(|| env::var("JWKS_URL").expect("JWKS_URL not set"));
async fn fetch_jwks() -> Result<Value, reqwest::Error> { async fn fetch_jwks() -> Result<Value, JwtValidationError> {
let jwks = reqwest::get(JWKS_URL.as_str()) let jwks = reqwest::get(JWKS_URL.as_str())
.await? .await?
.json::<Value>() .json::<Value>()
@@ -24,7 +26,7 @@ async fn fetch_jwks() -> Result<Value, reqwest::Error> {
Ok(jwks) Ok(jwks)
} }
pub async fn refresh_jwks() -> Result<Value, reqwest::Error> { pub async fn refresh_jwks() -> Result<Value, JwtValidationError> {
let jwks = fetch_jwks().await?; let jwks = fetch_jwks().await?;
let mut write = JWK_CACHE.write().await; let mut write = JWK_CACHE.write().await;
@@ -37,7 +39,7 @@ pub async fn refresh_jwks() -> Result<Value, reqwest::Error> {
Ok(jwks) Ok(jwks)
} }
pub async fn get_jwks() -> Result<Value, reqwest::Error> { pub async fn get_jwks() -> Result<Value, JwtValidationError> {
let ttl = Duration::from_secs(3600); // 1 hour let ttl = Duration::from_secs(3600); // 1 hour
{ {
+4
View File
@@ -0,0 +1,4 @@
pub mod claims;
pub mod errors;
pub mod jwks;
pub mod validator;
+70
View File
@@ -0,0 +1,70 @@
use super::claims::KeycloakClaims;
use super::errors::JwtValidationError;
use jsonwebtoken::{Algorithm, DecodingKey, Validation, decode, decode_header};
use once_cell::sync::Lazy;
use serde_json::Value;
use std::env;
static ISSUER: Lazy<String> = Lazy::new(|| env::var("ISSUER").expect("ISSUER not set"));
fn validate_token(token: &str, jwks: &Value) -> Result<KeycloakClaims, JwtValidationError> {
// 1. Decode header
let header = decode_header(token).map_err(|_| JwtValidationError::InvalidHeader)?;
let kid = header.kid.ok_or(JwtValidationError::MissingKid)?;
// 2. Find matching key
let keys = jwks["keys"]
.as_array()
.ok_or(JwtValidationError::InvalidJwks)?;
let key = keys
.iter()
.find(|k| k["kid"] == kid)
.ok_or(JwtValidationError::InvalidDecodingKey)?;
// 3. Extract RSA components
let n = key["n"]
.as_str()
.ok_or(JwtValidationError::MissingModulus)?;
let e = key["e"]
.as_str()
.ok_or(JwtValidationError::MissingExponent)?;
let decoding_key =
DecodingKey::from_rsa_components(n, e).map_err(|_| JwtValidationError::JwkNotFound)?;
// 4. Setup validation rules
let mut validation = Validation::new(Algorithm::RS256);
validation.set_issuer(&[ISSUER.as_str()]);
validation.validate_exp = true;
validation.validate_aud = false;
// 5. Decode & verify
let token_data = decode::<KeycloakClaims>(token, &decoding_key, &validation)
.map_err(|_| JwtValidationError::TokenValidationFailed)?;
Ok(token_data.claims)
}
pub async fn authenticate_jwt(token: &str) -> Result<KeycloakClaims, JwtValidationError> {
let jwks = super::jwks::get_jwks()
.await
.map_err(|_| JwtValidationError::JwksFetchFailed)?;
match validate_token(token, &jwks) {
Ok(claims) => Ok(claims),
Err(_) => {
// one retry with refresh
let fresh = super::jwks::refresh_jwks()
.await
.map_err(|_| JwtValidationError::JwksRefreshFailed)?;
validate_token(token, &fresh).map_err(|_| JwtValidationError::InvalidToken)
}
}
}
+2
View File
@@ -0,0 +1,2 @@
pub mod keycloak;
pub mod ollama;
+355
View File
@@ -0,0 +1,355 @@
use crate::providers::ollama;
use crate::providers::ollama::errors::LlmError;
use futures::StreamExt;
use reqwest::Client;
#[derive(Clone)]
pub struct OllamaProvider {
pub client: Client,
pub base_url: String,
}
impl OllamaProvider {
pub fn new(base_url: impl Into<String>) -> Self {
Self {
client: Client::new(),
base_url: base_url.into(),
}
}
// ── private helpers ──────────────────────────────────────────────────────
async fn model_exists(&self, model: &str) -> Result<bool, LlmError> {
let url = format!("{}/api/tags", self.base_url);
let res = self
.client
.get(url)
.send()
.await?
.json::<ollama::types::OllamaModels>()
.await?;
Ok(res.models.iter().any(|m| m.name == model))
}
fn has_user_message(&self, messages: &[ollama::types::OllamaMessage]) -> bool {
messages
.iter()
.any(|m| matches!(m.role, ollama::types::OllamaRole::User))
}
// pub fn validate_keep_alive(&self, s: &str) -> Result<(), LlmError> {
// let s = s.trim();
// if s == "-1" || s.parse::<u64>().is_ok() {
// return Ok(());
// }
// let split = s
// .find(|c: char| c.is_alphabetic())
// .ok_or_else(|| LlmError::InvalidKeepAlive(s.to_string()))?;
// let (num, unit) = s.split_at(split);
// num.parse::<u64>()
// .map_err(|_| LlmError::InvalidKeepAlive(s.to_string()))?;
// match unit {
// "s" | "m" | "h" => Ok(()),
// _ => Err(LlmError::InvalidKeepAlive(s.to_string())),
// }
// }
// // ── public endpoints ─────────────────────────────────────────────────────
pub async fn list_models(&self) -> Result<super::types::OllamaModels, LlmError> {
let url = format!("{}/api/tags", self.base_url);
let res = self
.client
.get(url)
.send()
.await?
.json::<ollama::types::OllamaModels>()
.await?;
Ok(res)
}
pub async fn completions(
&self,
body: &super::types::OllamaGenerateRequest,
) -> Result<super::types::OllamaGenerateResponse, LlmError> {
let url = format!("{}/api/generate", self.base_url);
if body.prompt.is_empty() {
return Err(LlmError::MissingPrompt);
}
if body.model.is_empty() {
return Err(LlmError::MissingModel);
}
let exists = self.model_exists(&body.model).await?;
if !exists {
return Err(LlmError::ModelNotFound(body.model.clone()));
}
let options = body.options.clone();
let payload = ollama::types::OllamaGenerateRequest {
model: body.model.clone(),
prompt: body.prompt.clone(),
stream: false,
keep_alive: body.keep_alive.clone(),
options,
};
dbg!(&payload);
let res = self
.client
.post(url)
.json(&payload)
.send()
.await?
.json::<ollama::types::OllamaGenerateResponse>()
.await?;
Ok(res)
}
pub async fn completions_stream(
&self,
body: &super::types::OllamaGenerateRequest,
) -> Result<ollama::types::OllamaGenerateResponseStream, LlmError> {
let url = format!("{}/api/generate", self.base_url);
if body.prompt.is_empty() {
return Err(LlmError::MissingPrompt);
}
if body.model.is_empty() {
return Err(LlmError::MissingModel);
}
let exists = self.model_exists(&body.model.to_string()).await?;
if !exists {
return Err(LlmError::ModelNotFound(body.model.to_string()));
}
let payload = ollama::types::OllamaGenerateRequest {
model: body.model.clone(),
prompt: body.prompt.clone(),
stream: true,
keep_alive: body.keep_alive.clone(),
options: body.options.clone(),
};
let byte_stream = self
.client
.post(url)
.json(&payload)
.send()
.await?
.bytes_stream();
let stream = byte_stream.flat_map(|chunk_result| {
let mut out: Vec<Result<super::types::OllamaGenerateStreamEvent, LlmError>> =
Vec::new();
let chunk = match chunk_result {
Ok(b) => b,
Err(e) => {
out.push(Err(LlmError::Http(e)));
return futures::stream::iter(out);
}
};
for line in chunk.split(|&b| b == b'\n') {
if line.is_empty() {
continue;
}
let parsed: ollama::types::OllamaGenerateResponse =
match serde_json::from_slice(line) {
Ok(v) => v,
Err(_) => continue,
};
if !parsed.response.is_empty() && !parsed.done {
out.push(Ok(super::types::OllamaGenerateStreamEvent::Token(
parsed.response.clone(),
)));
}
if parsed.done {
out.push(Ok(super::types::OllamaGenerateStreamEvent::Final(parsed)));
return futures::stream::iter(out);
}
}
futures::stream::iter(out)
});
Ok(Box::pin(stream))
}
pub async fn chat_completions(
&self,
body: &super::types::OllamaChatRequest,
) -> Result<super::types::OllamaChatResponse, LlmError> {
let url = format!("{}/api/chat", self.base_url);
if body.messages.is_empty() {
return Err(LlmError::MissingMessages);
}
let ollama_messages: Vec<ollama::types::OllamaMessage> = body.messages.clone();
if !self.has_user_message(&ollama_messages) {
return Err(LlmError::MissingMessages);
}
if body.model.is_empty() {
return Err(LlmError::MissingModel);
}
let exists = self.model_exists(&body.model).await?;
if !exists {
return Err(LlmError::ModelNotFound(body.model.clone()));
}
let options = body.options.clone();
let payload = ollama::types::OllamaChatRequest {
model: body.model.clone(),
messages: ollama_messages,
stream: false,
options,
keep_alive: body.keep_alive.clone(),
};
let res = self
.client
.post(url)
.json(&payload)
.send()
.await?
.json::<ollama::types::OllamaChatResponse>()
.await?;
Ok(res)
}
pub async fn chat_completions_stream(
&self,
body: &super::types::OllamaChatRequest,
) -> Result<ollama::types::OllamaChatResponseStream, LlmError> {
let url = format!("{}/api/chat", self.base_url);
if body.messages.is_empty() {
return Err(LlmError::MissingMessages);
}
if body.model.is_empty() {
return Err(LlmError::MissingModel);
}
let exists = self.model_exists(&body.model).await?;
if !exists {
return Err(LlmError::ModelNotFound(body.model.clone()));
}
let ollama_messages: Vec<ollama::types::OllamaMessage> = body.messages.clone();
if !self.has_user_message(&ollama_messages) {
return Err(LlmError::MissingMessages);
}
let payload = ollama::types::OllamaChatRequest {
model: body.model.clone(),
messages: body.messages.clone().into_iter().collect(),
stream: true,
options: body.options.clone(),
keep_alive: body.keep_alive.clone(),
};
let byte_stream = self
.client
.post(url)
.json(&payload)
.send()
.await?
.bytes_stream();
let stream = byte_stream.flat_map(|chunk_result| {
let mut events: Vec<Result<super::types::OllamaChatStreamEvent, LlmError>> = Vec::new();
let chunk = match chunk_result {
Ok(b) => b,
Err(e) => {
tracing::debug!("Error: {:?}", e);
events.push(Err(LlmError::Http(e)));
return futures::stream::iter(events);
}
};
let text = match std::str::from_utf8(&chunk) {
Ok(v) => v,
Err(e) => {
tracing::debug!("Invalid UTF8 from Ollama: {:?}", e);
return futures::stream::iter(events);
}
};
for line in text.lines() {
if line.trim().is_empty() {
continue;
}
let parsed: ollama::types::OllamaChatStreamResponse =
match serde_json::from_str(line) {
Ok(v) => v,
Err(e) => {
tracing::debug!("Failed parsing Ollama line {:?}: {:?}", line, e);
continue;
}
};
tracing::debug!("Parsed: {:?}", parsed);
// End of generation
if parsed.done {
println!("final: {:?}", &parsed);
events.push(Ok(super::types::OllamaChatStreamEvent::Final(
super::types::OllamaChatResponse {
model: parsed.model,
created_at: parsed.created_at,
message: parsed.message,
done_reason: parsed
.done_reason
.unwrap_or(super::types::OllamaFinishReason::Error),
total_duration: parsed.total_duration.unwrap_or(0),
load_duration: parsed.load_duration.unwrap_or(0),
prompt_eval_count: parsed.prompt_eval_count.unwrap_or(0),
eval_count: parsed.eval_count.unwrap_or(0),
},
)));
break;
}
// Normal generated token
if !parsed.message.content.is_empty() {
events.push(Ok(super::types::OllamaChatStreamEvent::Token(parsed)));
}
}
futures::stream::iter(events)
});
Ok(Box::pin(stream))
}
}
+23
View File
@@ -0,0 +1,23 @@
use thiserror::Error;
#[derive(Debug, Error)]
pub enum LlmError {
#[error("prompt is required and cannot be empty")]
MissingPrompt,
#[error("model is required and cannot be empty")]
MissingModel,
#[error("messages must be a non-empty array containing at least one user message")]
MissingMessages,
#[error("model '{0}' is not available — run `ollama pull {0}` first")]
ModelNotFound(String),
// #[error(
// "invalid keep_alive format '{0}' — expected <number><unit> (e.g. 30s, 10m, 2h), a plain integer (seconds), or -1"
// )]
// InvalidKeepAlive(String),
#[error(transparent)]
Http(#[from] reqwest::Error),
}
+3
View File
@@ -0,0 +1,3 @@
pub mod client;
pub mod errors;
pub mod types;
+161
View File
@@ -0,0 +1,161 @@
use crate::providers::ollama::errors;
use futures::Stream;
use serde::{Deserialize, Serialize};
use std::pin::Pin;
// ------ Models ------
#[derive(Debug, Deserialize)]
pub struct OllamaModels {
pub models: Vec<OllamaModel>,
}
#[derive(Debug, Deserialize)]
pub struct OllamaModel {
pub name: String,
pub details: Option<OllamaModelDetails>,
pub size: Option<u64>,
pub digest: Option<String>,
pub modified_at: Option<String>,
}
#[derive(Debug, Deserialize)]
pub struct OllamaModelDetails {
pub family: Option<String>,
pub parameter_size: Option<String>,
pub quantization_level: Option<String>,
}
// ------ Message ------
#[derive(Debug, Serialize, Deserialize, Clone)]
pub struct OllamaMessage {
pub role: OllamaRole,
pub content: String,
}
#[derive(Debug, Serialize, Deserialize, PartialEq, Eq, Clone)]
#[serde(rename_all = "lowercase")]
pub enum OllamaRole {
System,
User,
Assistant,
}
// ------ Shared ------
#[derive(Debug, Default, Clone, Serialize)]
pub struct OllamaOptions {
pub seed: Option<i64>,
pub temperature: Option<f32>,
pub top_p: Option<f32>,
pub top_k: Option<u32>,
pub stop: Option<Vec<String>>,
pub num_ctx: Option<u32>,
pub num_predict: Option<u32>,
}
// ------ Completion ------
#[derive(Debug, Serialize)]
pub struct OllamaGenerateRequest {
pub model: String,
pub prompt: String,
pub stream: bool,
pub keep_alive: String,
pub options: Option<OllamaOptions>,
}
#[derive(Debug, Deserialize)]
pub struct OllamaGenerateResponse {
pub model: String,
pub created_at: String,
pub response: String,
pub done: bool,
pub done_reason: Option<String>,
pub total_duration: Option<u64>,
pub load_duration: Option<u64>,
pub prompt_eval_count: Option<u32>,
pub eval_count: Option<u32>,
}
#[derive(Debug)]
pub enum OllamaGenerateStreamEvent {
Token(String),
Final(OllamaGenerateResponse),
}
pub type OllamaGenerateResponseStream = Pin<
Box<
dyn Stream<Item = Result<super::types::OllamaGenerateStreamEvent, errors::LlmError>> + Send,
>,
>;
// ------ Chat ------
#[derive(Debug, Serialize)]
pub struct OllamaChatRequest {
pub model: String,
pub messages: Vec<OllamaMessage>,
pub stream: bool,
pub keep_alive: String,
pub options: Option<OllamaOptions>,
}
#[derive(Debug, Deserialize)]
pub struct OllamaChatResponse {
pub model: String,
pub created_at: String,
pub message: OllamaMessage,
pub done_reason: OllamaFinishReason,
pub total_duration: u64,
pub load_duration: u64,
pub prompt_eval_count: u32,
pub eval_count: u32,
}
#[derive(Debug, Deserialize)]
pub struct OllamaChatStreamResponse {
pub model: String,
pub created_at: String,
pub message: OllamaMessage,
pub done: bool,
pub done_reason: Option<OllamaFinishReason>,
pub total_duration: Option<u64>,
pub load_duration: Option<u64>,
pub prompt_eval_count: Option<u32>,
pub eval_count: Option<u32>,
}
#[derive(Debug, Deserialize, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum OllamaFinishReason {
Stop,
Length,
Error,
}
#[derive(Debug)]
pub enum OllamaChatStreamEvent {
Token(OllamaChatStreamResponse),
Final(OllamaChatResponse),
}
pub type OllamaChatResponseStream = Pin<
Box<dyn Stream<Item = Result<super::types::OllamaChatStreamEvent, errors::LlmError>> + Send>,
>;
+100
View File
@@ -0,0 +1,100 @@
use crate::core;
use crate::services::errors::ServiceError;
use base64::{Engine as _, engine::general_purpose};
use rand::RngCore;
use rand::rngs::OsRng;
use sha2::{Digest, Sha256};
use sqlx::PgPool;
#[derive(Clone)]
pub struct AuthService {
postgres: PgPool,
}
fn generate_api_key() -> String {
let mut bytes = [0u8; 32];
OsRng.fill_bytes(&mut bytes);
general_purpose::URL_SAFE_NO_PAD.encode(bytes)
}
pub fn hash_key(key: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(key.as_bytes());
hasher
.finalize()
.iter()
.map(|b| format!("{:02x}", b))
.collect()
}
impl AuthService {
pub fn new(postgres: PgPool) -> Self {
Self { postgres }
}
// ------ API Key ------
pub async fn create_api_key(
&self,
payload: core::auth::api_key::CreateApiKeyRequest,
) -> Result<String, ServiceError> {
let raw_key = generate_api_key();
let key_hash = hash_key(&raw_key);
let key = crate::databases::postgres::api_key::types::CreateApiKey {
key_hash,
name: payload.name,
user_id: payload.user_id,
roles: payload.roles.into_iter().map(Into::into).collect(),
};
crate::databases::postgres::api_key::queries::create(&self.postgres, key).await?;
Ok(raw_key)
}
pub async fn validate_api_key(
&self,
key: &str,
) -> Result<core::auth::api_key::AuthContext, ServiceError> {
let key_hash: String = hash_key(key);
let auth =
crate::databases::postgres::api_key::queries::validate(&self.postgres, &key_hash)
.await?;
Ok(auth.into())
}
pub async fn update_last_access_api_key(
&self,
key_id: &uuid::Uuid,
) -> Result<(), ServiceError> {
crate::databases::postgres::api_key::queries::update_last_access(&self.postgres, key_id)
.await?;
Ok(())
}
// ------ JWT ------
pub async fn create_user(&self, user_id: &uuid::Uuid) -> Result<(), ServiceError> {
crate::databases::postgres::user_activity::queries::upsert_user_activity(
&self.postgres,
user_id,
)
.await?;
Ok(())
}
pub async fn validate_jwt(
&self,
token: &str,
) -> Result<core::auth::jwt::JwtClaims, ServiceError> {
let auth = crate::providers::keycloak::validator::authenticate_jwt(token).await?;
Ok(auth.into())
}
}
+357
View File
@@ -0,0 +1,357 @@
use crate::core;
use crate::core::llm;
use crate::core::llm::completions::{
CompletionResult, CompletionResultNoStream, CompletionStreamEvent,
};
use crate::providers::ollama::types::OllamaChatStreamEvent;
use crate::providers::{ollama::client::OllamaProvider, ollama::types::OllamaGenerateStreamEvent};
use crate::services::errors::ServiceError;
use super::ConversationService;
use async_stream::try_stream;
use futures::StreamExt;
use std::boxed::Box;
use std::sync::Arc;
use tokio::sync::Mutex;
use uuid::Uuid;
#[derive(Clone)]
pub struct ChatService {
ollama: OllamaProvider,
conversation: ConversationService,
}
impl ChatService {
pub fn new(ollama: OllamaProvider, conversation: ConversationService) -> Self {
Self {
ollama,
conversation,
}
}
pub async fn list_models(&self) -> Result<llm::models::Models, ServiceError> {
let models = self.ollama.list_models().await?;
Ok(models.into())
}
pub async fn load_model(
&self,
body: crate::core::llm::models::LoadModelRequest,
) -> Result<crate::core::llm::models::LoadModelResponse, ServiceError> {
let b = crate::providers::ollama::types::OllamaGenerateRequest {
model: body.model.clone(),
prompt: "load".to_string(),
stream: false,
keep_alive: body.keep_alive,
options: None,
};
self.ollama.completions(&b).await?;
Ok(crate::core::llm::models::LoadModelResponse { model: body.model })
}
pub async fn unload_model(
&self,
body: crate::core::llm::models::UnloadModelRequest,
) -> Result<crate::core::llm::models::UnloadModelResponse, ServiceError> {
let b = crate::providers::ollama::types::OllamaGenerateRequest {
model: body.model.clone(),
prompt: "unload".to_string(),
stream: false,
keep_alive: "0s".to_string(),
options: None,
};
self.ollama.completions(&b).await?;
Ok(crate::core::llm::models::UnloadModelResponse { model: body.model })
}
pub async fn complete(
&self,
body: core::llm::completions::CompletionRequest,
) -> Result<core::llm::completions::CompletionResult, ServiceError> {
let stream = body.options.stream;
let request: crate::providers::ollama::types::OllamaGenerateRequest = body.into();
if stream {
let ollama_stream = self.ollama.completions_stream(&request).await?;
let mapped = ollama_stream.map(|item| {
item.map(|event| match event {
OllamaGenerateStreamEvent::Token(tok) => CompletionStreamEvent::Token(tok),
OllamaGenerateStreamEvent::Final(resp) => {
CompletionStreamEvent::Final(CompletionResultNoStream {
id: Uuid::new_v4(),
created_at: resp.created_at,
model: resp.model,
text: resp.response,
prompt_tokens: resp.prompt_eval_count.unwrap_or(0),
completion_tokens: resp.eval_count.unwrap_or(0),
done_reason: resp.done_reason,
total_duration: resp.total_duration,
load_duration: resp.load_duration,
})
}
})
});
Ok(CompletionResult::Stream(Box::pin(mapped)))
} else {
let response = self.ollama.completions(&request).await?;
let enriched = core::llm::completions::CompletionResultNoStream {
id: Uuid::new_v4(),
created_at: response.created_at,
model: response.model,
text: response.response,
prompt_tokens: response.prompt_eval_count.unwrap_or(0),
completion_tokens: response.eval_count.unwrap_or(0),
done_reason: response.done_reason,
total_duration: response.total_duration,
load_duration: response.load_duration,
};
Ok(core::llm::completions::CompletionResult::NoStream(enriched))
}
}
pub async fn chat_complete(
&self,
body: core::llm::chat::ChatCompletionRequest,
auth: &core::auth::Auth,
) -> Result<core::llm::chat::ChatCompletionResult, ServiceError> {
let conversation_id = self
.resolve_conversation_with_title(auth.user_id(), &body)
.await?;
tracing::debug!(
"Received conversation_id={:?}, parent_id={:?}",
conversation_id,
body.parent_id
);
let user_msg_id = self
.conversation
.log_user_message(
auth.user_id(),
conversation_id,
body.parent_id,
&body.message.content,
None,
)
.await?;
let stream = body.options.stream;
let history = self
.build_chat_history(auth.user_id(), conversation_id, body.options.context_depth)
.await?;
let mut request: crate::providers::ollama::types::OllamaChatRequest = body.into();
request.messages = history.into_iter().map(Into::into).collect();
if stream {
let mut ollama_stream = Box::pin(self.ollama.chat_completions_stream(&request).await?);
let conversation_svc = self.conversation.clone();
let user_id = auth.user_id();
let out = try_stream! {
let accumulated = Arc::new(Mutex::new(String::new()));
// Pull first item to get real model/created_at for Start.
let first = ollama_stream.next().await;
let Some(first) = first else {
Err(ServiceError::Internal(
"provider stream ended before producing any events".to_string(),
))?;
return;
};
let first = first?; // propagates provider error via `?` inside try_stream!
let (model, created_at) = match &first {
OllamaChatStreamEvent::Token(tok) => (tok.model.clone(), tok.created_at.clone()),
OllamaChatStreamEvent::Final(resp) => (resp.model.clone(), resp.created_at.clone()),
};
yield core::llm::chat::ChatCompletionStreamEvent::Start {
model,
conversation_id,
message_id: user_msg_id,
created_at,
};
// Helper closure-like inline handling so we don't duplicate match logic;
// process `first`, then continue draining the rest of the stream.
let mut pending = Some(first);
loop {
let item = match pending.take() {
Some(item) => item,
None => match ollama_stream.next().await {
Some(res) => res?,
None => break,
},
};
match item {
OllamaChatStreamEvent::Token(tok) => {
accumulated.lock().await.push_str(&tok.message.content);
yield core::llm::chat::ChatCompletionStreamEvent::Token {
content: tok.message.content,
created_at: tok.created_at,
id: Uuid::new_v4(),
};
}
OllamaChatStreamEvent::Final(mut resp) => {
let content = accumulated.lock().await.clone();
resp.message.content = content.clone();
let assistant_message_id = conversation_svc
.log_assistant_message(
user_id,
conversation_id,
user_msg_id,
&content,
resp.eval_count,
)
.await?;
yield core::llm::chat::ChatCompletionStreamEvent::Final(
core::llm::chat::ChatCompletionResultNoStream {
id: assistant_message_id,
conversation_id,
created_at: resp.created_at,
model: resp.model,
message: resp.message.into(),
prompt_tokens: resp.prompt_eval_count,
completion_tokens: resp.eval_count,
finish_reason: resp.done_reason.into(),
total_duration: resp.total_duration,
load_duration: resp.load_duration,
},
);
}
}
}
};
Ok(core::llm::chat::ChatCompletionResult::Stream(Box::pin(out)))
} else {
let response = self.ollama.chat_completions(&request).await?;
let assistant_message_id = self
.conversation
.log_assistant_message(
auth.user_id(),
conversation_id,
user_msg_id,
&response.message.content,
response.eval_count,
)
.await?;
self.conversation
.update_message_tokens(auth.user_id(), user_msg_id, response.prompt_eval_count)
.await?;
let enriched = core::llm::chat::ChatCompletionResultNoStream {
id: assistant_message_id,
conversation_id,
created_at: response.created_at,
model: response.model,
message: response.message.into(),
prompt_tokens: response.prompt_eval_count,
completion_tokens: response.eval_count,
finish_reason: response.done_reason.into(),
total_duration: response.total_duration,
load_duration: response.load_duration,
};
Ok(core::llm::chat::ChatCompletionResult::NoStream(enriched))
}
}
// ------ Helpers ------
async fn resolve_conversation_with_title(
&self,
user_id: Uuid,
body: &core::llm::chat::ChatCompletionRequest,
) -> Result<Uuid, ServiceError> {
let conversation_id = match self
.conversation
.get_or_create_conversation(user_id, body.conversation_id)
.await?
{
core::databases::conversations::ConversationResult::Existing(id) => id,
core::databases::conversations::ConversationResult::Created(id) => {
let last_message = body.message.content.as_str();
let prompt = format!(
"Generate a title using next message in maximum 6 words: {}",
last_message
);
let title_result = self
.complete(core::llm::completions::CompletionRequest {
model: body.model.clone(),
prompt,
options: core::llm::completions::CompletionOptions {
stream: false,
..Default::default()
},
})
.await?;
if let CompletionResult::NoStream(t) = title_result {
self.conversation
.set_conversation_title(user_id, id, &t.text)
.await?;
}
id
}
};
Ok(conversation_id)
}
async fn build_chat_history(
&self,
auth_user_id: Uuid,
conversation_id: Uuid,
context_depth: u32,
) -> Result<Vec<crate::core::llm::chat::Message>, ServiceError> {
let messages = self
.conversation
.get_messages_entries(
auth_user_id,
conversation_id,
crate::core::databases::conversations::CursorPage {
limit: context_depth,
before: None,
},
)
.await?;
let mut history: Vec<_> = messages
.messages
.into_iter()
.map(crate::core::llm::chat::Message::from_summary)
.collect();
history.reverse();
Ok(history)
}
}
+176
View File
@@ -0,0 +1,176 @@
use crate::core;
use crate::databases::postgres;
use crate::services::errors::ServiceError;
use sqlx::PgPool;
use uuid::Uuid;
#[derive(Clone)]
pub struct ConversationService {
postgres: PgPool,
}
impl ConversationService {
pub fn new(postgres: PgPool) -> Self {
Self { postgres }
}
pub async fn get_conversations_entries(
&self,
user_id: Uuid,
pointer: crate::core::databases::conversations::CursorPage,
) -> Result<core::databases::conversations::ConversationList, ServiceError> {
let limit = pointer.limit;
let conversations = postgres::chat::queries::get_conversations_entries(
&self.postgres,
user_id,
limit,
pointer.before,
)
.await?;
let has_more = conversations.len() == limit as usize;
Ok(crate::core::databases::conversations::ConversationList {
conversations: conversations.into_iter().map(Into::into).collect(),
has_more,
})
}
pub async fn get_messages_entries(
&self,
user_id: Uuid,
conversation_id: Uuid,
pointer: crate::core::databases::conversations::CursorPage,
) -> Result<core::databases::conversations::MessageList, ServiceError> {
let limit = pointer.limit;
let mut messages = postgres::chat::queries::get_conversation_messages(
&self.postgres,
user_id,
conversation_id,
limit + 1,
pointer.before,
)
.await?;
let has_more = messages.len() == limit as usize;
if has_more {
messages.pop();
}
messages.reverse();
Ok(crate::core::databases::conversations::MessageList {
messages: messages.into_iter().map(Into::into).collect(),
has_more,
})
}
pub async fn get_or_create_conversation(
&self,
user_id: Uuid,
conversation_id: Option<Uuid>,
) -> Result<core::databases::conversations::ConversationResult, ServiceError> {
let state = postgres::chat::queries::get_or_create_conversation(
&self.postgres,
conversation_id,
user_id,
)
.await?;
Ok(state.into())
}
pub async fn set_conversation_title(
&self,
user_id: Uuid,
conversation_id: Uuid,
title: &str,
) -> Result<(), ServiceError> {
postgres::chat::queries::set_conversation_title(
&self.postgres,
conversation_id,
user_id,
title,
)
.await?;
Ok(())
}
pub async fn log_message(
&self,
user_id: Uuid,
conversation_id: Uuid,
parent_id: Option<Uuid>,
role: core::llm::ChatRole,
content: &str,
tokens: Option<u32>,
) -> Result<Uuid, ServiceError> {
let id = postgres::chat::queries::insert_message(
&self.postgres,
user_id,
conversation_id,
parent_id,
role.into(),
content,
tokens,
)
.await?;
Ok(id)
}
pub async fn log_user_message(
&self,
user_id: Uuid,
conversation_id: Uuid,
parent_id: Option<Uuid>,
content: &str,
tokens: Option<u32>,
) -> Result<Uuid, ServiceError> {
self.log_message(
user_id,
conversation_id,
parent_id,
crate::core::llm::ChatRole::User,
content,
tokens,
)
.await
}
pub async fn log_assistant_message(
&self,
user_id: Uuid,
conversation_id: Uuid,
parent_id: Uuid,
content: &str,
tokens: u32,
) -> Result<Uuid, ServiceError> {
self.log_message(
user_id,
conversation_id,
Some(parent_id),
crate::core::llm::ChatRole::Assistant,
content,
Some(tokens),
)
.await
}
pub async fn update_message_tokens(
&self,
user_id: Uuid,
message_id: Uuid,
tokens: u32,
) -> Result<(), ServiceError> {
postgres::chat::queries::update_message_tokens(&self.postgres, user_id, message_id, tokens)
.await?;
Ok(())
}
}
+29
View File
@@ -0,0 +1,29 @@
use crate::databases::errors::DbError;
use crate::providers::keycloak::errors::JwtValidationError;
use crate::providers::ollama::errors::LlmError;
#[derive(Debug)]
pub enum ServiceError {
Db(DbError),
Llm(LlmError),
Auth(JwtValidationError),
Internal(String),
}
impl From<DbError> for ServiceError {
fn from(e: DbError) -> Self {
ServiceError::Db(e)
}
}
impl From<LlmError> for ServiceError {
fn from(e: LlmError) -> Self {
ServiceError::Llm(e)
}
}
impl From<JwtValidationError> for ServiceError {
fn from(e: JwtValidationError) -> Self {
ServiceError::Auth(e)
}
}
+8
View File
@@ -0,0 +1,8 @@
pub mod auth_service;
pub mod chat_service;
pub mod conversation_service;
pub mod errors;
pub use auth_service::AuthService;
pub use chat_service::ChatService;
pub use conversation_service::ConversationService;
+419
View File
@@ -0,0 +1,419 @@
// use serde_json::json;
// use wiremock::matchers::{method, path};
// use wiremock::{Mock, MockServer, ResponseTemplate};
// use chat::dto::api;
// use chat::providers::ollama::client::OllamaProvider;
// use chat::providers::ollama::errors::OllamaError;
// // ── helpers ──────────────────────────────────────────────────────────────────
// async fn setup() -> (MockServer, OllamaProvider) {
// let server = MockServer::start().await;
// let provider = OllamaProvider::new(server.uri());
// (server, provider)
// }
// fn models_response(names: &[&str]) -> serde_json::Value {
// json!({
// "models": names.iter().map(|n| json!({ "name": n })).collect::<Vec<_>>()
// })
// }
// // ── list_models ───────────────────────────────────────────────────────────────
// #[tokio::test]
// async fn test_list_models_ok() {
// let (server, provider) = setup().await;
// Mock::given(method("GET"))
// .and(path("/api/tags"))
// .respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
// .mount(&server)
// .await;
// let res = provider.list_models().await.unwrap();
// assert_eq!(res.models[0].name, "llama3");
// }
// // ── completions ───────────────────────────────────────────────────────────────
// #[tokio::test]
// async fn test_completions_ok() {
// let (server, provider) = setup().await;
// Mock::given(method("GET"))
// .and(path("/api/tags"))
// .respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
// .mount(&server)
// .await;
// Mock::given(method("POST"))
// .and(path("/api/generate"))
// .respond_with(ResponseTemplate::new(200).set_body_json(json!({
// "model": "llama3",
// "response": "I am a helpful assistant.",
// "done": true,
// "prompt_eval_count": 10,
// "eval_count": 8,
// })))
// .mount(&server)
// .await;
// let req = api::CompletionRequest {
// base: api::BaseLLMRequest {
// model: "llama3".to_string(),
// ..Default::default()
// },
// prompt: "Hello".to_string(),
// };
// let res = provider.completions(&req).await.unwrap();
// assert_eq!(res.object, api::CompletionObject::TextCompletion);
// assert_eq!(res.choices.len(), 1);
// assert_eq!(res.choices[0].text, "I am a helpful assistant.");
// assert_eq!(res.choices[0].finish_reason, api::FinishReason::Stop);
// assert_eq!(res.usage.prompt_tokens, 10);
// assert_eq!(res.usage.completion_tokens, 8);
// assert_eq!(res.usage.total_tokens, 18);
// }
// #[tokio::test]
// async fn test_completions_missing_prompt() {
// let (server, provider) = setup().await;
// Mock::given(method("GET"))
// .and(path("/api/tags"))
// .respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
// .mount(&server)
// .await;
// let req = api::CompletionRequest {
// base: api::BaseLLMRequest {
// model: "llama3".to_string(),
// ..Default::default()
// },
// prompt: "".to_string(),
// };
// let err = provider.completions(&req).await.unwrap_err();
// assert!(matches!(err, OllamaError::MissingPrompt));
// }
// #[tokio::test]
// async fn test_completions_empty_prompt() {
// let (server, provider) = setup().await;
// Mock::given(method("GET"))
// .and(path("/api/tags"))
// .respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
// .mount(&server)
// .await;
// let req = api::CompletionRequest {
// base: api::BaseLLMRequest {
// model: "llama3".to_string(),
// ..Default::default()
// },
// prompt: " ".to_string(),
// };
// let err = provider.completions(&req).await.unwrap_err();
// assert!(matches!(err, OllamaError::MissingPrompt));
// }
// #[tokio::test]
// async fn test_completions_model_not_found() {
// let (server, provider) = setup().await;
// Mock::given(method("GET"))
// .and(path("/api/tags"))
// .respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
// .mount(&server)
// .await;
// let req = api::CompletionRequest {
// base: api::BaseLLMRequest {
// model: "gpt-4".to_string(),
// ..Default::default()
// },
// prompt: "hello".to_string(),
// };
// let err = provider.completions(&req).await.unwrap_err();
// assert!(matches!(err, OllamaError::ModelNotFound(_)));
// }
// // ── chat_completions ──────────────────────────────────────────────────────────
// #[tokio::test]
// async fn test_chat_completions_ok() {
// let (server, provider) = setup().await;
// Mock::given(method("GET"))
// .and(path("/api/tags"))
// .respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
// .mount(&server)
// .await;
// Mock::given(method("POST"))
// .and(path("/api/chat"))
// .respond_with(ResponseTemplate::new(200).set_body_json(json!({
// "model": "llama3",
// "message": { "role": "assistant", "content": "4." },
// "done": true,
// "prompt_eval_count": 5,
// "eval_count": 2
// })))
// .mount(&server)
// .await;
// let req = api::ChatRequest {
// base: api::BaseLLMRequest {
// model: "llama3".to_string(),
// ..Default::default()
// },
// messages: vec![api::Message {
// role: api::Role::User,
// content: "What is 2+2?".to_string(),
// }],
// conversation_id: None,
// parent_id: None,
// };
// let res = provider.chat_completions(&req).await.unwrap();
// assert_eq!(res.object, "chat.completion");
// assert_eq!(res.choices.len(), 1);
// assert_eq!(res.choices[0].message.role, api::Role::Assistant);
// assert_eq!(res.choices[0].message.content, "4.");
// assert_eq!(res.choices[0].finish_reason, api::FinishReason::Stop);
// let usage = res.usage.unwrap();
// assert_eq!(usage.prompt_tokens, 5);
// assert_eq!(usage.completion_tokens, 2);
// assert_eq!(usage.total_tokens, 7);
// }
// #[tokio::test]
// async fn test_chat_completions_missing_messages() {
// let (server, provider) = setup().await;
// Mock::given(method("GET"))
// .and(path("/api/tags"))
// .respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
// .mount(&server)
// .await;
// let req = api::ChatRequest {
// base: api::BaseLLMRequest {
// model: "llama3".to_string(),
// ..Default::default()
// },
// messages: vec![],
// conversation_id: None,
// parent_id: None,
// };
// let err = provider.chat_completions(&req).await.unwrap_err();
// assert!(matches!(err, OllamaError::MissingMessages));
// }
// #[tokio::test]
// async fn test_chat_completions_no_user_message() {
// let (server, provider) = setup().await;
// Mock::given(method("GET"))
// .and(path("/api/tags"))
// .respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
// .mount(&server)
// .await;
// let req = api::ChatRequest {
// base: api::BaseLLMRequest {
// model: "llama3".to_string(),
// ..Default::default()
// },
// messages: vec![api::Message {
// role: api::Role::System,
// content: "be helpful".to_string(),
// }],
// conversation_id: None,
// parent_id: None,
// };
// let err = provider.chat_completions(&req).await.unwrap_err();
// assert!(matches!(err, OllamaError::MissingMessages));
// }
// #[tokio::test]
// async fn test_chat_completions_model_not_found() {
// let (server, provider) = setup().await;
// Mock::given(method("GET"))
// .and(path("/api/tags"))
// .respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
// .mount(&server)
// .await;
// let req = api::ChatRequest {
// base: api::BaseLLMRequest {
// model: "gpt-4".to_string(),
// ..Default::default()
// },
// messages: vec![api::Message {
// role: api::Role::User,
// content: "hi".to_string(),
// }],
// conversation_id: None,
// parent_id: None,
// };
// let err = provider.chat_completions(&req).await.unwrap_err();
// assert!(matches!(err, OllamaError::ModelNotFound(_)));
// }
// // // ── load_model ────────────────────────────────────────────────────────────────
// #[tokio::test]
// async fn test_load_model_ok() {
// let (server, provider) = setup().await;
// Mock::given(method("GET"))
// .and(path("/api/tags"))
// .respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
// .mount(&server)
// .await;
// Mock::given(method("POST"))
// .and(path("/api/generate"))
// .respond_with(ResponseTemplate::new(200).set_body_json(json!({
// "model": "llama3",
// "response": "ok",
// "done": true,
// })))
// .mount(&server)
// .await;
// let res = provider.load_model("llama3", Some("10m")).await.unwrap();
// assert_eq!(res.model, "llama3");
// assert_eq!(res.status, "loaded");
// assert_eq!(res.keep_alive, "10m");
// }
// #[tokio::test]
// async fn test_load_model_not_found() {
// let (server, provider) = setup().await;
// Mock::given(method("GET"))
// .and(path("/api/tags"))
// .respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
// .mount(&server)
// .await;
// let err = provider.load_model("gpt-4", Some("10m")).await.unwrap_err();
// assert!(matches!(err, OllamaError::ModelNotFound(_)));
// }
// #[tokio::test]
// async fn test_load_model_invalid_keep_alive() {
// let (server, provider) = setup().await;
// Mock::given(method("GET"))
// .and(path("/api/tags"))
// .respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
// .mount(&server)
// .await;
// let err = provider
// .load_model("llama3", Some("10x"))
// .await
// .unwrap_err();
// assert!(matches!(err, OllamaError::InvalidKeepAlive(_)));
// }
// #[tokio::test]
// async fn test_load_model_keep_alive_plain_integer() {
// let (server, provider) = setup().await;
// Mock::given(method("GET"))
// .and(path("/api/tags"))
// .respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
// .mount(&server)
// .await;
// Mock::given(method("POST"))
// .and(path("/api/generate"))
// .respond_with(ResponseTemplate::new(200).set_body_json(json!({
// "model": "llama3",
// "done": true,
// })))
// .mount(&server)
// .await;
// let res = provider.load_model("llama3", Some("3600")).await.unwrap();
// assert_eq!(res.status, "loaded");
// assert_eq!(res.keep_alive, "3600");
// let res = provider.load_model("llama3", Some("-1")).await.unwrap();
// assert_eq!(res.status, "loaded");
// assert_eq!(res.keep_alive, "-1");
// }
// // // ── unload_model ──────────────────────────────────────────────────────────────
// #[tokio::test]
// async fn test_unload_model_ok() {
// let (server, provider) = setup().await;
// Mock::given(method("GET"))
// .and(path("/api/tags"))
// .respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
// .mount(&server)
// .await;
// Mock::given(method("POST"))
// .and(path("/api/generate"))
// .respond_with(ResponseTemplate::new(200).set_body_json(json!({
// "model": "llama3",
// "response": "ok",
// "done": true,
// })))
// .mount(&server)
// .await;
// let res = provider.unload_model("llama3").await.unwrap();
// assert_eq!(res.model, "llama3");
// assert_eq!(res.status, "unloaded");
// }
// #[tokio::test]
// async fn test_unload_model_not_found() {
// let (server, provider) = setup().await;
// Mock::given(method("GET"))
// .and(path("/api/tags"))
// .respond_with(ResponseTemplate::new(200).set_body_json(models_response(&["llama3"])))
// .mount(&server)
// .await;
// let err = provider.unload_model("gpt-4").await.unwrap_err();
// assert!(matches!(err, OllamaError::ModelNotFound(_)));
// }