From 28352d8bdb107c5a733bdc629f65bb857888c462 Mon Sep 17 00:00:00 2001 From: LucasDLTG Date: Tue, 2 Jun 2026 18:30:39 +0200 Subject: [PATCH] reafctor: all code without stream --- ...8f643a724b729f73f8308ae0b3a70d87dc5de.json | 17 - ...b268f3910e4e3266f029f17f460c9f5362dd4.json | 51 ++ ...dd427bf246f4eabd3e335077f47a7aa7882e4.json | 14 - ...9cb4004548fc2efc7cbcb056e11e2534eb441.json | 34 - ...7af232e6a9e9010c947422f0497ff76b06c58.json | 34 + ...af61028fd0ae9f34c13504950be6742f107f.json} | 19 +- ...2af653bcf39aa1bb4d8add80ed069421ee859.json | 13 +- ...b9bcbafad49d7e1df83cb15a1b86bbdeb4fdf.json | 14 + Cargo.lock | 47 +- Cargo.toml | 14 +- readme.md | 13 + src/api/app.rs | 57 ++ src/api/docs.rs | 55 ++ src/api/errors.rs | 212 +++-- src/api/middlewares/auth.rs | 162 ++++ src/{ => api}/middlewares/mod.rs | 0 src/api/mod.rs | 6 + src/{ => api}/routes/mod.rs | 0 src/api/routes/v1/apikey.rs | 25 + src/api/routes/v1/chat.rs | 337 +++++++ src/{ => api}/routes/v1/mod.rs | 31 +- src/api/routes/v1/models.rs | 138 +++ src/api/state/app_state.rs | 12 + src/api/state/mod.rs | 4 + src/{dto/api.rs => api/types.rs} | 186 ++-- src/core/auth/api_key.rs | 27 + src/core/auth/jwt.rs | 25 + src/core/auth/mod.rs | 33 + src/core/databases/conversations.rs | 44 + src/core/databases/mod.rs | 1 + src/core/llm/chat.rs | 79 ++ src/core/llm/completions.rs | 58 ++ src/core/llm/mod.rs | 6 + src/core/llm/models.rs | 36 + src/core/llm/role.rs | 8 + src/core/mod.rs | 3 + src/databases/{postgres => }/errors.rs | 3 + src/databases/mod.rs | 2 + src/databases/postgres/api_key/mod.rs | 1 + src/databases/postgres/api_key/queries.rs | 44 +- src/databases/postgres/api_key/types.rs | 23 + src/databases/postgres/chat/queries.rs | 24 +- src/databases/postgres/chat/types.rs | 9 +- src/databases/postgres/mod.rs | 3 +- src/databases/postgres/pool.rs | 2 +- src/databases/postgres/user/queries.rs | 20 - .../postgres/{user => user_activity}/mod.rs | 0 .../postgres/user_activity/queries.rs | 21 + src/docs.rs | 55 -- src/dto/mod.rs | 2 - src/dto/ollama.rs | 80 -- src/lib.rs | 7 +- src/main.rs | 72 +- src/mappers/api_to_core.rs | 66 ++ src/mappers/core_to_api.rs | 101 +++ src/mappers/core_to_database.rs | 30 + src/mappers/core_to_ollama.rs | 67 ++ src/mappers/database_to_core.rs | 69 ++ src/mappers/keycloak_to_core.rs | 30 + src/mappers/mod.rs | 7 + src/mappers/ollama_to_core.rs | 50 ++ src/middlewares/auth/apikey.rs | 15 - src/middlewares/auth/keycloak.rs | 148 ---- src/middlewares/auth/middleware.rs | 227 ----- src/middlewares/auth/mod.rs | 5 - src/providers/keycloak/claims.rs | 23 + src/providers/keycloak/errors.rs | 40 + src/providers/keycloak/jwks.rs | 59 ++ src/providers/keycloak/mod.rs | 4 + src/providers/keycloak/validator.rs | 64 ++ src/providers/mod.rs | 1 + src/providers/ollama/client.rs | 414 ++++----- src/providers/ollama/errors.rs | 45 +- src/providers/ollama/mapper.rs | 88 -- src/providers/ollama/mod.rs | 2 +- src/providers/ollama/types.rs | 137 +++ src/routes/v1/apikey.rs | 50 -- src/routes/v1/chat.rs | 504 ----------- src/routes/v1/models.rs | 134 --- src/routes/v1/openapi.rs | 8 - src/services/auth_service.rs | 100 +++ src/services/chat_service.rs | 317 +++++++ src/services/conversation_service.rs | 132 +++ src/services/errors.rs | 27 + src/services/mod.rs | 8 + src/state/app_state.rs | 9 - src/state/mod.rs | 1 - src/utils/crypto.rs | 11 - src/utils/mod.rs | 1 - tests/ollama_provider.rs | 838 +++++++++--------- 90 files changed, 3581 insertions(+), 2434 deletions(-) delete mode 100644 .sqlx/query-14b38b0f3fca61893dc6e036dc68f643a724b729f73f8308ae0b3a70d87dc5de.json create mode 100644 .sqlx/query-52a8daad9aa79a9c391bfa52b92b268f3910e4e3266f029f17f460c9f5362dd4.json delete mode 100644 .sqlx/query-70117c16fc3efeebbcb40838d13dd427bf246f4eabd3e335077f47a7aa7882e4.json delete mode 100644 .sqlx/query-931f64701ca90a159a54dab717b9cb4004548fc2efc7cbcb056e11e2534eb441.json create mode 100644 .sqlx/query-ac5047be40cf3c686d3a3f877fc7af232e6a9e9010c947422f0497ff76b06c58.json rename .sqlx/{query-aa2722be6d9aec100c73bf65ef14ce3e979afe85a6d9d2b6e0005b9848efbd77.json => query-ac5b689f154241e193cff087a552af61028fd0ae9f34c13504950be6742f107f.json} (54%) create mode 100644 .sqlx/query-f389b37680c97698323c81cb7eeb9bcbafad49d7e1df83cb15a1b86bbdeb4fdf.json create mode 100644 src/api/app.rs create mode 100644 src/api/docs.rs create mode 100644 src/api/middlewares/auth.rs rename src/{ => api}/middlewares/mod.rs (100%) rename src/{ => api}/routes/mod.rs (100%) create mode 100644 src/api/routes/v1/apikey.rs create mode 100644 src/api/routes/v1/chat.rs rename src/{ => api}/routes/v1/mod.rs (55%) create mode 100644 src/api/routes/v1/models.rs create mode 100644 src/api/state/app_state.rs create mode 100644 src/api/state/mod.rs rename src/{dto/api.rs => api/types.rs} (63%) create mode 100644 src/core/auth/api_key.rs create mode 100644 src/core/auth/jwt.rs create mode 100644 src/core/auth/mod.rs create mode 100644 src/core/databases/conversations.rs create mode 100644 src/core/databases/mod.rs create mode 100644 src/core/llm/chat.rs create mode 100644 src/core/llm/completions.rs create mode 100644 src/core/llm/mod.rs create mode 100644 src/core/llm/models.rs create mode 100644 src/core/llm/role.rs create mode 100644 src/core/mod.rs rename src/databases/{postgres => }/errors.rs (82%) create mode 100644 src/databases/postgres/api_key/types.rs delete mode 100644 src/databases/postgres/user/queries.rs rename src/databases/postgres/{user => user_activity}/mod.rs (100%) create mode 100644 src/databases/postgres/user_activity/queries.rs delete mode 100644 src/docs.rs delete mode 100644 src/dto/mod.rs delete mode 100644 src/dto/ollama.rs create mode 100644 src/mappers/api_to_core.rs create mode 100644 src/mappers/core_to_api.rs create mode 100644 src/mappers/core_to_database.rs create mode 100644 src/mappers/core_to_ollama.rs create mode 100644 src/mappers/database_to_core.rs create mode 100644 src/mappers/keycloak_to_core.rs create mode 100644 src/mappers/mod.rs create mode 100644 src/mappers/ollama_to_core.rs delete mode 100644 src/middlewares/auth/apikey.rs delete mode 100644 src/middlewares/auth/keycloak.rs delete mode 100644 src/middlewares/auth/middleware.rs delete mode 100644 src/middlewares/auth/mod.rs create mode 100644 src/providers/keycloak/claims.rs create mode 100644 src/providers/keycloak/errors.rs create mode 100644 src/providers/keycloak/jwks.rs create mode 100644 src/providers/keycloak/mod.rs create mode 100644 src/providers/keycloak/validator.rs delete mode 100644 src/providers/ollama/mapper.rs create mode 100644 src/providers/ollama/types.rs delete mode 100644 src/routes/v1/apikey.rs delete mode 100644 src/routes/v1/chat.rs delete mode 100644 src/routes/v1/models.rs delete mode 100644 src/routes/v1/openapi.rs create mode 100644 src/services/auth_service.rs create mode 100644 src/services/chat_service.rs create mode 100644 src/services/conversation_service.rs create mode 100644 src/services/errors.rs create mode 100644 src/services/mod.rs delete mode 100644 src/state/app_state.rs delete mode 100644 src/state/mod.rs delete mode 100644 src/utils/crypto.rs delete mode 100644 src/utils/mod.rs diff --git a/.sqlx/query-14b38b0f3fca61893dc6e036dc68f643a724b729f73f8308ae0b3a70d87dc5de.json b/.sqlx/query-14b38b0f3fca61893dc6e036dc68f643a724b729f73f8308ae0b3a70d87dc5de.json deleted file mode 100644 index 5b23f83..0000000 --- a/.sqlx/query-14b38b0f3fca61893dc6e036dc68f643a724b729f73f8308ae0b3a70d87dc5de.json +++ /dev/null @@ -1,17 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n INSERT INTO auth.api_key (key_hash, name, created_by, scopes)\n VALUES ($1, $2, $3, $4)\n ", - "describe": { - "columns": [], - "parameters": { - "Left": [ - "Text", - "Text", - "Uuid", - "TextArray" - ] - }, - "nullable": [] - }, - "hash": "14b38b0f3fca61893dc6e036dc68f643a724b729f73f8308ae0b3a70d87dc5de" -} diff --git a/.sqlx/query-52a8daad9aa79a9c391bfa52b92b268f3910e4e3266f029f17f460c9f5362dd4.json b/.sqlx/query-52a8daad9aa79a9c391bfa52b92b268f3910e4e3266f029f17f460c9f5362dd4.json new file mode 100644 index 0000000..4a3cc55 --- /dev/null +++ b/.sqlx/query-52a8daad9aa79a9c391bfa52b92b268f3910e4e3266f029f17f460c9f5362dd4.json @@ -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\"\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", + "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" +} diff --git a/.sqlx/query-70117c16fc3efeebbcb40838d13dd427bf246f4eabd3e335077f47a7aa7882e4.json b/.sqlx/query-70117c16fc3efeebbcb40838d13dd427bf246f4eabd3e335077f47a7aa7882e4.json deleted file mode 100644 index 51039f7..0000000 --- a/.sqlx/query-70117c16fc3efeebbcb40838d13dd427bf246f4eabd3e335077f47a7aa7882e4.json +++ /dev/null @@ -1,14 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n INSERT INTO auth.app_user (id)\n VALUES ($1)\n ON CONFLICT (id) DO NOTHING\n ", - "describe": { - "columns": [], - "parameters": { - "Left": [ - "Uuid" - ] - }, - "nullable": [] - }, - "hash": "70117c16fc3efeebbcb40838d13dd427bf246f4eabd3e335077f47a7aa7882e4" -} diff --git a/.sqlx/query-931f64701ca90a159a54dab717b9cb4004548fc2efc7cbcb056e11e2534eb441.json b/.sqlx/query-931f64701ca90a159a54dab717b9cb4004548fc2efc7cbcb056e11e2534eb441.json deleted file mode 100644 index 6975824..0000000 --- a/.sqlx/query-931f64701ca90a159a54dab717b9cb4004548fc2efc7cbcb056e11e2534eb441.json +++ /dev/null @@ -1,34 +0,0 @@ -{ - "db_name": "PostgreSQL", - "query": "\n SELECT u.id AS user_id, ak.id as key_id, ak.scopes as roles\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": "key_id", - "type_info": "Uuid" - }, - { - "ordinal": 2, - "name": "roles", - "type_info": "TextArray" - } - ], - "parameters": { - "Left": [ - "Text" - ] - }, - "nullable": [ - false, - false, - false - ] - }, - "hash": "931f64701ca90a159a54dab717b9cb4004548fc2efc7cbcb056e11e2534eb441" -} diff --git a/.sqlx/query-ac5047be40cf3c686d3a3f877fc7af232e6a9e9010c947422f0497ff76b06c58.json b/.sqlx/query-ac5047be40cf3c686d3a3f877fc7af232e6a9e9010c947422f0497ff76b06c58.json new file mode 100644 index 0000000..07a03a5 --- /dev/null +++ b/.sqlx/query-ac5047be40cf3c686d3a3f877fc7af232e6a9e9010c947422f0497ff76b06c58.json @@ -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" +} diff --git a/.sqlx/query-aa2722be6d9aec100c73bf65ef14ce3e979afe85a6d9d2b6e0005b9848efbd77.json b/.sqlx/query-ac5b689f154241e193cff087a552af61028fd0ae9f34c13504950be6742f107f.json similarity index 54% rename from .sqlx/query-aa2722be6d9aec100c73bf65ef14ce3e979afe85a6d9d2b6e0005b9848efbd77.json rename to .sqlx/query-ac5b689f154241e193cff087a552af61028fd0ae9f34c13504950be6742f107f.json index 4b004bf..6a81a31 100644 --- a/.sqlx/query-aa2722be6d9aec100c73bf65ef14ce3e979afe85a6d9d2b6e0005b9848efbd77.json +++ b/.sqlx/query-ac5b689f154241e193cff087a552af61028fd0ae9f34c13504950be6742f107f.json @@ -1,6 +1,6 @@ { "db_name": "PostgreSQL", - "query": "\n SELECT id, parent_id, role, 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 ASC\n LIMIT $3\n ", + "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": [ { @@ -15,8 +15,19 @@ }, { "ordinal": 2, - "name": "role", - "type_info": "Text" + "name": "role: MessageRole", + "type_info": { + "Custom": { + "name": "chat.role", + "kind": { + "Enum": [ + "user", + "assistant", + "system" + ] + } + } + } }, { "ordinal": 3, @@ -50,5 +61,5 @@ true ] }, - "hash": "aa2722be6d9aec100c73bf65ef14ce3e979afe85a6d9d2b6e0005b9848efbd77" + "hash": "ac5b689f154241e193cff087a552af61028fd0ae9f34c13504950be6742f107f" } diff --git a/.sqlx/query-c98b93484092fc57098f4c6eac92af653bcf39aa1bb4d8add80ed069421ee859.json b/.sqlx/query-c98b93484092fc57098f4c6eac92af653bcf39aa1bb4d8add80ed069421ee859.json index 7ae1020..6a0a226 100644 --- a/.sqlx/query-c98b93484092fc57098f4c6eac92af653bcf39aa1bb4d8add80ed069421ee859.json +++ b/.sqlx/query-c98b93484092fc57098f4c6eac92af653bcf39aa1bb4d8add80ed069421ee859.json @@ -13,7 +13,18 @@ "Left": [ "Uuid", "Uuid", - "Text", + { + "Custom": { + "name": "chat.role", + "kind": { + "Enum": [ + "user", + "assistant", + "system" + ] + } + } + }, "Text", "Int4" ] diff --git a/.sqlx/query-f389b37680c97698323c81cb7eeb9bcbafad49d7e1df83cb15a1b86bbdeb4fdf.json b/.sqlx/query-f389b37680c97698323c81cb7eeb9bcbafad49d7e1df83cb15a1b86bbdeb4fdf.json new file mode 100644 index 0000000..ce5245a --- /dev/null +++ b/.sqlx/query-f389b37680c97698323c81cb7eeb9bcbafad49d7e1df83cb15a1b86bbdeb4fdf.json @@ -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" +} diff --git a/Cargo.lock b/Cargo.lock index f5600d7..ce1fff3 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -93,6 +93,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "31b698c5f9a010f6573133b09e0de5408834d0c82f8d7475a89fc1867a71cd90" dependencies = [ "axum-core", + "axum-macros", "bytes", "form_urlencoded", "futures-util", @@ -138,6 +139,17 @@ dependencies = [ "tracing", ] +[[package]] +name = "axum-macros" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7aa268c23bfbbd2c4363b9cd302a4f504fb2a9dfe7e3451d66f35dd392e20aca" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "base64" version = "0.22.1" @@ -1165,9 +1177,9 @@ dependencies = [ [[package]] name = "jsonwebtoken" -version = "10.3.0" +version = "10.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0529410abe238729a60b108898784df8984c87f6054c9c4fcacc47e4803c1ce1" +checksum = "eba32bfb4ffdeaca3e34431072faf01745c9b26d25504aa7a6cf5684334fc4fc" dependencies = [ "aws-lc-rs", "base64", @@ -1178,6 +1190,7 @@ dependencies = [ "serde_json", "signature", "simple_asn1", + "zeroize", ] [[package]] @@ -1718,9 +1731,9 @@ checksum = "dc897dd8d9e8bd1ed8cdad82b5966c3e0ecae09fb1907d58efaa013543185d0a" [[package]] name = "reqwest" -version = "0.13.3" +version = "0.13.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "62e0021ea2c22aed41653bc7e1419abb2c97e038ff2c33d0e1309e49a97deec0" +checksum = "219c5811de6525e5416c7d5d53bb656d3afdbc6c5af816e0802bcfa42dbdc1c3" dependencies = [ "base64", "bytes", @@ -1972,9 +1985,9 @@ dependencies = [ [[package]] name = "serde_json" -version = "1.0.149" +version = "1.0.150" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "83fc039473c5595ace860d8c4fafa220ff474b3fc6bfdb4293327f1a37e94d86" +checksum = "e8014e44b4736ed0538adeecded0fce2a272f22dc9578a7eb6b2d9993c74cfb9" dependencies = [ "itoa", "memchr", @@ -2588,9 +2601,9 @@ dependencies = [ [[package]] name = "tower-http" -version = "0.6.10" +version = "0.6.11" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "68d6fdd9f81c2819c9a8b0e0cd91660e7746a8e6ea2ba7c6b2b057985f6bcb51" +checksum = "4cfcf7e2740e6fc6d4d688b4ef00650406bb94adf4731e43c096c3a19fe40840" dependencies = [ "bitflags", "bytes", @@ -2780,9 +2793,9 @@ dependencies = [ [[package]] name = "uuid" -version = "1.23.1" +version = "1.23.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ddd74a9687298c6858e9b88ec8935ec45d22e8fd5e6394fa1bd4e99a87789c76" +checksum = "d258b83ceec21034727ecee8c382cfa6c3e133699b0742c64571814fb420c9f7" dependencies = [ "getrandom 0.4.2", "js-sys", @@ -3569,6 +3582,20 @@ name = "zeroize" version = "1.8.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b97154e67e32c85465826e8bcc1c59429aaaf107c1e4a9e53c8d8ccd5eff88d0" +dependencies = [ + "zeroize_derive", +] + +[[package]] +name = "zeroize_derive" +version = "1.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85a5b4158499876c763cb03bc4e49185d3cccbabb15b33c627f7884f43db852e" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] [[package]] name = "zerotrie" diff --git a/Cargo.toml b/Cargo.toml index 4b4d1bb..38b7409 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -8,24 +8,24 @@ wiremock = "0.6" tokio = { version = "1.52.3", features = ["macros", "rt-multi-thread"] } [dependencies] -axum = "0.8.9" +axum = { version = "0.8.9", features = ["macros"] } utoipa = { version = "5.5.0", features = ["axum_extras", "uuid"] } tokio = { version = "1", features = ["full"] } serde = { version = "1", features = ["derive"] } -serde_json = "1" -jsonwebtoken = { version = "10.3.0", features = ["aws_lc_rs"] } -reqwest = { version = "0.13.3", features = ["json", "stream"] } +serde_json = "1.0.150" +jsonwebtoken = { version = "10.4.0", features = ["aws_lc_rs"] } +reqwest = { version = "0.13.4", features = ["json", "stream"] } once_cell = "1" dotenvy = "0.15" thiserror = "2.0.18" tokio-stream = "0.1" futures = "0.3" chrono = { version = "0.4.44", features = ["serde"] } -uuid = { version = "1.23.1", features = ["v4", "serde"] } -tower-http = { version = "0.6.10", features = ["cors"] } +uuid = { version = "1.23.2", features = ["v4", "serde"] } +tower-http = { version = "0.6.11", 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"] } +sqlx = { version = "0.8.6", features = ["runtime-tokio-rustls", "postgres", "uuid", "chrono", "macros"] } rand = "0.8" base64 = "0.22.1" sha2 = "0.11.0" \ No newline at end of file diff --git a/readme.md b/readme.md index 163d101..acf46d0 100644 --- a/readme.md +++ b/readme.md @@ -19,6 +19,17 @@ curl -s -X POST https://chat.iceberg.black/api/v1/chat/completions \ ] }' | 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. @@ -266,3 +277,5 @@ This project turns Ollama into: # TODO - open api doc for bearer token +- Unify check before sending to ollama payload +- load/unload model functions \ No newline at end of file diff --git a/src/api/app.rs b/src/api/app.rs new file mode 100644 index 0000000..94bcfaa --- /dev/null +++ b/src/api/app.rs @@ -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 = 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 api_key_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: api_key_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::().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) +} diff --git a/src/api/docs.rs b/src/api/docs.rs new file mode 100644 index 0000000..d91865d --- /dev/null +++ b/src/api/docs.rs @@ -0,0 +1,55 @@ +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::chat::completions, + // routes::v1::chat::chat_completions, + routes::v1::models::list_models, + // routes::v1::models::load_model, + // routes::v1::models::unload_model, + ), + components( + schemas( + api::errors::ErrorResponse, + api::types::ModelsResponse, + api::types::ModelInfo, + api::types::LoadModelResponse, + api::types::LoadModelRequest, + api::types::UnloadModelResponse, + api::types::LLMOptions, + api::types::CompletionRequest, + api::types::CompletionObject, + api::types::FinishReason, + api::types::CompletionResponse, + api::types::Choice, + api::types::Usage, + api::types::CompletionChunk, + api::types::ChatRequest, + api::types::Message, + api::types::Role, + api::types::ChatResponse, + api::types::ChatChoice, + api::types::ChatCompletionChunk, + api::types::ChatChunkChoice, + api::types::Delta, + ) + ), + tags( + (name = "chat", description = "Chat & completions"), + (name = "models", description = "Model management") + ) +)] +pub struct ApiDoc; diff --git a/src/api/errors.rs b/src/api/errors.rs index 554a504..470f048 100644 --- a/src/api/errors.rs +++ b/src/api/errors.rs @@ -1,114 +1,148 @@ -use crate::databases::postgres::errors::DbError; -use crate::providers::ollama::errors::OllamaError; +use crate::databases::errors::DbError; +use crate::providers::keycloak::errors::AuthError; +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 thiserror::Error; +use utoipa::ToSchema; -#[derive(Serialize)] +#[derive(Serialize, ToSchema)] pub struct ErrorResponse { + pub status: u16, pub error: String, pub code: String, } -pub enum ApiError { - Db(DbError), - Ollama(OllamaError), +pub struct ApiError { + pub status: StatusCode, + pub code: &'static str, + pub message: String, } -impl From for ApiError { - fn from(e: DbError) -> Self { - ApiError::Db(e) +#[derive(Debug, Error)] +pub enum AuthMiddlewareError { + #[error("invalid authorization format")] + InvalidAuthorizationFormat, + + #[error("authentication required")] + AuthenticationRequired, + + #[error("insufficient permissions")] + Forbidden, +} + +impl From 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 for ApiError { - fn from(e: OllamaError) -> Self { - ApiError::Ollama(e) +impl From 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 { + AuthError::InvalidToken | AuthError::TokenValidationFailed => Self { + status: StatusCode::UNAUTHORIZED, + code: "AUTH_INVALID_TOKEN", + message: "invalid or expired token".into(), + }, + AuthError::InvalidHeader + // | AuthError::InvalidHeaderDecode + | AuthError::MissingKid => Self { + status: StatusCode::UNAUTHORIZED, + code: "AUTH_INVALID_HEADER", + message: "invalid authorization header".into(), + }, + AuthError::JwkNotFound + | AuthError::InvalidJwks + | AuthError::MissingModulus + | AuthError::MissingExponent + | AuthError::InvalidDecodingKey => Self { + status: StatusCode::INTERNAL_SERVER_ERROR, + code: "AUTH_JWKS_ERROR", + message: "key validation error".into(), + }, + AuthError::JwksFetchFailed + | AuthError::JwksRefreshFailed + | AuthError::Reqwest(_) => Self { + status: StatusCode::SERVICE_UNAVAILABLE, + code: "AUTH_JWKS_FETCH", + message: "failed to fetch authorization keys".into(), + }, + }, + } } } impl IntoResponse for ApiError { fn into_response(self) -> Response { - let (status, body) = match self { - ApiError::Db(db_err) => match db_err { - DbError::Connection(_) => ( - StatusCode::INTERNAL_SERVER_ERROR, - ErrorResponse { - error: "database connection error".to_string(), - code: "DB_CONNECTION".to_string(), - }, - ), - DbError::Timeout => ( - StatusCode::REQUEST_TIMEOUT, - ErrorResponse { - error: "database timeout".to_string(), - code: "DB_TIMEOUT".to_string(), - }, - ), - DbError::NotFound => ( - StatusCode::NOT_FOUND, - ErrorResponse { - error: "not found".to_string(), - code: "DB_NOT_FOUND".to_string(), - }, - ), - }, - - ApiError::Ollama(err) => match err { - OllamaError::MissingPrompt => ( - StatusCode::BAD_REQUEST, - ErrorResponse { - error: "prompt is required and cannot be empty".to_string(), - code: "OLLAMA_MISSING_PROMPT".to_string(), - }, - ), - OllamaError::MissingModel => ( - StatusCode::BAD_REQUEST, - ErrorResponse { - error: "model is required and cannot be empty".to_string(), - code: "OLLAMA_MISSING_MODEL".to_string(), - }, - ), - OllamaError::ModelNotFound(m) => ( - StatusCode::UNPROCESSABLE_ENTITY, - ErrorResponse { - error: format!("model '{m}' is not available"), - code: "OLLAMA_MODEL_NOT_FOUND".to_string(), - }, - ), - OllamaError::MissingKeepAlive => ( - StatusCode::BAD_REQUEST, - ErrorResponse { - error: "keep_alive is required".to_string(), - code: "OLLAMA_MISSING_KEEPALIVE".to_string(), - }, - ), - OllamaError::InvalidKeepAlive(v) => ( - StatusCode::BAD_REQUEST, - ErrorResponse { - error: format!("invalid keep_alive '{v}'"), - code: "OLLAMA_INVALID_KEEPALIVE".to_string(), - }, - ), - OllamaError::MissingMessages => ( - StatusCode::BAD_REQUEST, - ErrorResponse { - error: "messages array required".to_string(), - code: "OLLAMA_MISSING_MESSAGES".to_string(), - }, - ), - OllamaError::Http(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - ErrorResponse { - error: e.to_string(), - code: "OLLAMA_HTTP_ERROR".to_string(), - }, - ), - }, + let body = ErrorResponse { + status: self.status.as_u16(), + error: self.message, + code: self.code.to_string(), }; - (status, Json(body)).into_response() + (self.status, Json(body)).into_response() } } diff --git a/src/api/middlewares/auth.rs b/src/api/middlewares/auth.rs new file mode 100644 index 0000000..28438ff --- /dev/null +++ b/src/api/middlewares/auth.rs @@ -0,0 +1,162 @@ +use crate::api::errors::{ApiError, 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), +} + +pub async fn auth_middleware( + State(state): State, + 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) => { + let err = match (jwt_err, api_key_err) { + // Both headers absent + (JwtError::MissingHeader, ApiKeyError::MissingHeader) => { + AuthMiddlewareError::AuthenticationRequired + } + // JWT header present but malformed — surface it, api key result irrelevant + (JwtError::InvalidFormat, _) => AuthMiddlewareError::InvalidAuthorizationFormat, + // JWT service failure — api key header was missing, so JWT was the intended method + (JwtError::Service(e), ApiKeyError::MissingHeader) => { + return e.into_response(); + } + // Both services failed + (JwtError::Service(e), ApiKeyError::Service(_)) => { + return e.into_response(); + } + // JWT missing, api key service failed + (JwtError::MissingHeader, ApiKeyError::Service(e)) => { + return e.into_response(); + } + }; + + return ApiError::from(err).into_response(); + } + }, + }; + + match handle_auth(&state, req, next, auth).await { + Ok(response) => response, + Err(err) => err.into_response(), + } +} + +async fn try_jwt( + state: &SharedState, + headers: &HeaderMap, +) -> Result { + 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 { + 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 handle_auth( + state: &SharedState, + mut request: Request, + next: Next, + auth: crate::core::auth::Auth, +) -> Result { + 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?; + } + } + + request.extensions_mut().insert(auth); + + Ok(next.run(request).await) +} + +// ── 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 { + let auth = request + .extensions() + .get::() + .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) +} diff --git a/src/middlewares/mod.rs b/src/api/middlewares/mod.rs similarity index 100% rename from src/middlewares/mod.rs rename to src/api/middlewares/mod.rs diff --git a/src/api/mod.rs b/src/api/mod.rs index 629e98f..4257309 100644 --- a/src/api/mod.rs +++ b/src/api/mod.rs @@ -1 +1,7 @@ +pub mod app; +pub mod docs; pub mod errors; +pub mod middlewares; +pub mod routes; +pub mod state; +pub mod types; diff --git a/src/routes/mod.rs b/src/api/routes/mod.rs similarity index 100% rename from src/routes/mod.rs rename to src/api/routes/mod.rs diff --git a/src/api/routes/v1/apikey.rs b/src/api/routes/v1/apikey.rs new file mode 100644 index 0000000..3dcf7d0 --- /dev/null +++ b/src/api/routes/v1/apikey.rs @@ -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, + Extension(claims): Extension, + Json(body): Json, +) -> Result, 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 })) +} diff --git a/src/api/routes/v1/chat.rs b/src/api/routes/v1/chat.rs new file mode 100644 index 0000000..f0c76e5 --- /dev/null +++ b/src/api/routes/v1/chat.rs @@ -0,0 +1,337 @@ +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( + post, + path = "/completions", + tag = "chat", + request_body( + content = api::types::CompletionRequest, + 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::CompletionResponse, + 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, + Json(body): Json, +) -> Result { + tracing::debug!("Received /completion with body {:?}", body); + + let response = state.chat_service.complete(body.into()).await?; + + match response { + CompletionResult::NoStream(res) => Ok(Json::< + crate::core::llm::completions::CompletionResultNoStream, + >(res) + .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 = "/chat/completions", + tag = "chat", + request_body( + content = api::types::ChatRequest, + 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::ChatResponse, + 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, + Extension(auth): Extension, + Json(body): Json, +) -> Result { + 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::(res).into_response()) + } + crate::core::llm::chat::ChatCompletionResult::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()) + } + } +} + +// Conversation retrieveing +pub async fn get_conversations( + State(state): State, + Extension(auth): Extension, + Query(params): Query, +) -> Result, 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, + Extension(auth): Extension, + Path(conversation_id): Path, + Query(params): Query, +) -> Result, 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())) +} + +// async fn handle_stream( +// state: AppState, +// auth: Auth, +// body: api::ChatRequest, +// conversation_id: Option, +// ) -> Result { +// // Handle anonymous (API key) path early — no DB logging +// let Some(conv_id) = conversation_id else { +// let stream = state.ollama.chat_completions_stream(&body).await?; + +// let plain_stream = stream.map( +// |item| -> Result { +// match item { +// Ok(chunk) => Ok( +// Event::default().data(serde_json::to_string(&chunk).unwrap_or_default()) +// ), +// Err(e) => Err(e), +// } +// }, +// ); +// return Ok(Sse::new(plain_stream) +// .keep_alive(KeepAlive::default()) +// .into_response()); +// }; + +// // From here conv_id is a plain Uuid — all variables stay in scope +// let user_msg_id = log_user_message( +// &state.postgres, +// auth.user_id(), +// conv_id, +// body.parent_id, +// body.messages +// .last() +// .map(|m| m.content.as_str()) +// .unwrap_or(""), +// None, +// ) +// .await?; + +// let start_event = api::StreamEvent::Start(api::StartEventData { +// conversation_id: conv_id, +// created: chrono::Utc::now().timestamp() as u64, +// id: user_msg_id, +// }); + +// let (tx, rx) = tokio::sync::mpsc::channel::< +// Result, +// >(32); + +// // Send start event immediately, before Ollama is contacted +// let _ = tx +// .send(Ok(Event::default() +// .event("metadata") +// .data(serde_json::to_string(&start_event).unwrap()))) +// .await; + +// let pool = state.postgres.clone(); +// let user_id = auth.user_id(); + +// tokio::spawn(async move { +// // Ollama called inside spawn — start event already queued +// let stream = match state.ollama.chat_completions_stream(&body).await { +// Ok(s) => s, +// Err(e) => { +// let _ = tx.send(Err(e)).await; +// return; +// } +// }; + +// let mut stream = stream; +// let mut accumulated = String::new(); + +// while let Some(item) = futures::StreamExt::next(&mut stream).await { +// match item { +// Ok(chunk) => { +// let is_done = chunk.choices[0].finish_reason == Some(api::FinishReason::Stop); + +// if let Some(content) = chunk.choices[0].delta.content.as_ref() { +// accumulated.push_str(content); +// } + +// if is_done { +// let prompt_tokens = chunk.usage.as_ref().map(|u| u.prompt_tokens); +// let completion_tokens = chunk.usage.as_ref().map(|u| u.completion_tokens); + +// if let Some(pt) = prompt_tokens { +// let _ = update_message_tokens(&pool, user_id, user_msg_id, pt).await; +// } + +// let assistant_msg_id = log_assistant_message( +// &pool, +// user_id, +// conv_id, +// user_msg_id, +// &accumulated, +// completion_tokens, +// ) +// .await; + +// if let Ok(msg_id) = assistant_msg_id { +// let end_event = api::StreamEvent::End(api::EndEventData { +// usage: api::Usage { +// prompt_tokens: prompt_tokens.unwrap_or(0), +// completion_tokens: completion_tokens.unwrap_or(0), +// total_tokens: chunk +// .usage +// .as_ref() +// .map(|u| u.total_tokens) +// .unwrap_or(0), +// }, +// id: msg_id, +// created: chrono::Utc::now().timestamp() as u64, +// }); + +// let _ = tx +// .send(Ok(Event::default() +// .event("metadata") +// .data(serde_json::to_string(&end_event).unwrap()))) +// .await; +// } + +// break; +// } + +// let data = api::StreamEvent::Delta(chunk); +// let json = serde_json::to_string(&data).unwrap(); +// if tx.send(Ok(Event::default().data(json))).await.is_err() { +// break; +// } +// } +// Err(e) => { +// let _ = tx.send(Err(e)).await; +// break; +// } +// } +// } +// }); + +// Ok(Sse::new(tokio_stream::wrappers::ReceiverStream::new(rx)) +// .keep_alive(KeepAlive::default()) +// .into_response()) +// } diff --git a/src/routes/v1/mod.rs b/src/api/routes/v1/mod.rs similarity index 55% rename from src/routes/v1/mod.rs rename to src/api/routes/v1/mod.rs index 1fdce57..3cd966e 100644 --- a/src/routes/v1/mod.rs +++ b/src/api/routes/v1/mod.rs @@ -2,33 +2,36 @@ pub mod apikey; pub mod chat; pub mod models; -use crate::docs::ApiDoc; -use crate::middlewares::auth::{auth_middleware, middleware::require_roles}; -use crate::state::app_state::AppState; +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, routing::post}; +use axum::{ + Json, Router, middleware, + routing::{get, post}, +}; use utoipa::OpenApi; async fn openapi_json() -> Json { Json(ApiDoc::openapi()) } -fn public_router() -> Router { +fn public_router() -> Router { Router::new().route("/docs.json", get(openapi_json)) } -pub fn protected_router() -> Router { +pub fn protected_router() -> Router { Router::new() .route("/models", get(models::list_models)) .route("/completions", post(chat::completions)) .route("/chat/completions", post(chat::chat_completions)) .route("/models/{model}/load", post(models::load_model)) - .route("/models/{model}/unload", post(models::unload_model)) + // .route("/models/{model}/unload", post(models::unload_model)) .route( "/keys/generate", - post(apikey::create_api_key).route_layer(middleware::from_fn(|req, next| { - require_roles(req, next, None, Some("admin")) - })), + post(apikey::create_api_key) // Usage + .route_layer(role_guard!(Some("admin"), None)), ) .route("/conversations", get(chat::get_conversations)) .route( @@ -37,8 +40,12 @@ pub fn protected_router() -> Router { ) } -pub fn router(state: AppState) -> Router { +pub fn router(state: SharedState) -> Router { Router::new() .merge(public_router()) - .merge(protected_router().layer(middleware::from_fn_with_state(state, auth_middleware))) + .merge(protected_router().layer(middleware::from_fn_with_state( + state.clone(), + auth_middleware, + ))) + .with_state(state) } diff --git a/src/api/routes/v1/models.rs b/src/api/routes/v1/models.rs new file mode 100644 index 0000000..11cd32a --- /dev/null +++ b/src/api/routes/v1/models.rs @@ -0,0 +1,138 @@ +use crate::api::{self, state::SharedState}; + +use axum::{ + Json, + extract::{Path, State}, +}; + +#[utoipa::path( + get, + path = "/models", + tag = "models", + responses( + ( + status = 200, + description = "List of locally available Ollama models", + body = api::types::ModelsResponse, + content_type = "application/json", + ), + ( + status = 500, + description = "Internal server error (Ollama or network failure)", + body = api::errors::ErrorResponse, + example = json!({ "error": "connection refused" }) + ) + ) +)] +// #[axum::debug_handler] +pub async fn list_models( + State(state): State, +) -> Result, api::errors::ApiError> { + let models = state.chat_service.list_models().await?; + + Ok(Json(models.into())) +} + +#[utoipa::path( + post, + path = "/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::LoadModelRequest, + 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::LoadModelResponse, + 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, + Path(model): Path, + Json(body): Json, +) -> Result, 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::LoadModelResponse { + model: response.model, + keep_alive: body.keep_alive, + status: "loaded".to_string(), + })) +} + +// #[utoipa::path( +// delete, +// path = "/models/{model}/load", +// 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::UnloadModelResponse, +// 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, +// Path(model): Path, +// ) -> Result, (axum::http::StatusCode, String)> { +// let response = state +// .ollama +// .unload_model(&model) +// .await +// .map_err(into_http_response)?; + +// Ok(Json(response)) +// } diff --git a/src/api/state/app_state.rs b/src/api/state/app_state.rs new file mode 100644 index 0000000..1076f19 --- /dev/null +++ b/src/api/state/app_state.rs @@ -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; diff --git a/src/api/state/mod.rs b/src/api/state/mod.rs new file mode 100644 index 0000000..e6a16d2 --- /dev/null +++ b/src/api/state/mod.rs @@ -0,0 +1,4 @@ +pub mod app_state; + +pub use app_state::AppState; +pub use app_state::SharedState; diff --git a/src/dto/api.rs b/src/api/types.rs similarity index 63% rename from src/dto/api.rs rename to src/api/types.rs index a2c8306..39de1be 100644 --- a/src/dto/api.rs +++ b/src/api/types.rs @@ -1,31 +1,32 @@ use serde::{Deserialize, Serialize}; +use std::collections::HashMap; use utoipa::ToSchema; use uuid::Uuid; +// ------ Models ------ + #[derive(Debug, Serialize, ToSchema)] -pub struct ErrorResponse { - pub error: String, -} - -impl ErrorResponse { - pub fn new(msg: impl Into) -> Self { - Self { error: msg.into() } - } -} - -#[derive(Debug, Serialize, Deserialize, ToSchema)] -pub struct ModelsResponse { - pub models: Vec, -} - -#[derive(Debug, Serialize, Deserialize, ToSchema)] pub struct ModelInfo { pub name: String, pub family: Option, pub parameter_size: Option, - pub quantization: Option, + pub metadata: ModelMetadata, } +#[derive(Debug, Clone, Default, Serialize, ToSchema)] +pub struct ModelMetadata { + pub extra: HashMap, +} + +// ------ Endpoint: /models ------ + +#[derive(Debug, Serialize, ToSchema)] +pub struct ModelsResponse { + pub models: Vec, +} + +// ------ Endpoint: /models/{model}/load ------ + #[derive(Debug, Serialize, Deserialize, ToSchema)] pub struct LoadModelResponse { pub model: String, @@ -34,27 +35,28 @@ pub struct LoadModelResponse { } #[derive(Debug, Serialize, Deserialize, ToSchema)] -pub struct LoadModelBody { - pub keep_alive: Option, +pub struct LoadModelRequest { + pub keep_alive: String, } +// ------ Endpoint: /models/{model}/unload ------ + #[derive(Debug, Serialize, Deserialize, ToSchema)] pub struct UnloadModelResponse { pub model: String, pub status: String, } -#[derive(Debug, Deserialize, Serialize, ToSchema, Default)] -pub struct BaseLLMRequest { - pub model: String, +// ------ Completions ------ +#[derive(Debug, Deserialize, Serialize, ToSchema, Default)] +pub struct LLMOptions { #[serde(default)] pub stream: bool, pub temperature: Option, pub top_p: Option, - // Ollama-native pub top_k: Option, pub repeat_penalty: Option, pub seed: Option, @@ -69,20 +71,35 @@ pub struct BaseLLMRequest { pub context_depth: Option, } +// ------ Endpoint: /completions ------ + #[derive(Debug, Deserialize, Serialize, ToSchema)] pub struct CompletionRequest { #[serde(flatten)] - pub base: BaseLLMRequest, + pub options: LLMOptions, + pub model: String, pub prompt: String, } +#[derive(Debug, Serialize, Deserialize, ToSchema)] +pub struct CompletionResponse { + pub id: Uuid, + pub object: CompletionObject, + pub created: String, + pub model: String, + pub choices: Vec, + pub usage: Usage, +} + #[derive(Debug, Serialize, Deserialize, ToSchema, PartialEq, Eq)] #[serde(rename_all = "snake_case")] pub enum CompletionObject { TextCompletion, } +// ------ Endpoint: /chat/completions ------ + #[derive(Debug, Serialize, Deserialize, ToSchema, PartialEq, Eq)] #[serde(rename_all = "snake_case")] pub enum FinishReason { @@ -93,16 +110,6 @@ pub enum FinishReason { Error, } -#[derive(Debug, Serialize, Deserialize, ToSchema)] -pub struct CompletionResponse { - pub id: String, - pub object: CompletionObject, - pub created: u64, - pub model: String, - pub choices: Vec, - pub usage: Usage, -} - #[derive(Debug, Serialize, Deserialize, ToSchema)] pub struct Choice { pub text: String, @@ -127,13 +134,14 @@ pub struct CompletionChunk { #[derive(Debug, Deserialize, Serialize, ToSchema)] pub struct ChatRequest { #[serde(flatten)] - pub base: BaseLLMRequest, + pub base: LLMOptions, - pub messages: Vec, + pub model: String, + pub message: Message, // Non standard Open AI pub conversation_id: Option, - pub parent_id: Option, + pub parent_id: Option, // used when branching, regenerate, etc } #[derive(Debug, Deserialize, Serialize, ToSchema)] @@ -151,15 +159,15 @@ pub enum Role { } #[derive(Debug, Serialize, Deserialize, ToSchema)] -pub struct ChatCompletionResponse { +pub struct ChatResponse { pub id: String, pub object: String, pub created: u64, pub model: String, pub choices: Vec, - pub usage: Option, // optional (Ollama may not always provide) + pub usage: Usage, - pub conversation_id: Option, + pub conversation_id: Uuid, } #[derive(Debug, Serialize, Deserialize, ToSchema)] @@ -190,62 +198,96 @@ pub struct Delta { pub role: Option, } -#[derive(Debug, Serialize, Deserialize, ToSchema)] -pub struct StartEventData { - pub conversation_id: Uuid, - pub created: u64, - pub id: Uuid, -} +// #[derive(Debug, Serialize, Deserialize, ToSchema)] +// pub struct StartEventData { +// pub conversation_id: Uuid, +// pub created: u64, +// pub id: Uuid, +// } -#[derive(Debug, Serialize, Deserialize, ToSchema)] -pub struct EndEventData { - pub created: u64, - pub id: Uuid, - pub usage: Usage, -} +// #[derive(Debug, Serialize, Deserialize, ToSchema)] +// pub struct EndEventData { +// pub created: u64, +// pub id: Uuid, +// pub usage: Usage, +// } -#[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), +// #[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, + pub scopes: Vec, } #[derive(serde::Serialize)] pub struct CreateApiKeyResponse { - pub api_key: String, // ONLY returned once + pub api_key: String, } +// ------ Fetch database ------ +// --- Shared --- #[derive(Debug, Deserialize)] -pub struct ConversationQuery { - pub limit: Option, - pub before: Option>, // cursor +pub struct CursorPage { + pub limit: Option, + pub before: Option>, +} + +// --- Conversation --- + +#[derive(Debug, Serialize)] +pub struct ConversationSummary { + pub id: Uuid, + pub title: String, + pub created_at: chrono::DateTime, + pub updated_at: chrono::DateTime, } #[derive(Debug, Serialize)] pub struct ConversationListResponse { - pub conversations: Vec, + pub conversations: Vec, pub has_more: bool, } -#[derive(Debug, Deserialize)] -pub struct MessageQuery { - pub limit: Option, - pub before: Option>, +// --- Message --- + +#[derive(Debug, Serialize)] +pub enum ApiChatRole { + System, + User, + Assistant, +} + +#[derive(Debug, Serialize)] +pub struct MessageSummary { + pub id: Uuid, + pub parent_id: Option, + pub role: ApiChatRole, + pub content: String, + pub created_at: chrono::DateTime, + pub tokens: Option, } #[derive(Debug, Serialize)] pub struct MessageListResponse { - pub messages: Vec, + pub messages: Vec, pub has_more: bool, } diff --git a/src/core/auth/api_key.rs b/src/core/auth/api_key.rs new file mode 100644 index 0000000..2150918 --- /dev/null +++ b/src/core/auth/api_key.rs @@ -0,0 +1,27 @@ +use uuid::Uuid; + +#[derive(Debug)] +pub struct CreateApiKeyRequest { + pub user_id: Uuid, + pub name: String, + pub roles: Vec, +} + +#[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, +} + +impl AuthContext { + pub fn has_role(&self, role: &KeyRole) -> bool { + self.roles.iter().any(|r| r == role) + } +} diff --git a/src/core/auth/jwt.rs b/src/core/auth/jwt.rs new file mode 100644 index 0000000..f9bb1e8 --- /dev/null +++ b/src/core/auth/jwt.rs @@ -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, + pub _exp: usize, + pub _issuer: String, + + pub realm_roles: Vec, + pub client_roles: HashMap>, +} + +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)) + } +} diff --git a/src/core/auth/mod.rs b/src/core/auth/mod.rs new file mode 100644 index 0000000..4ac6434 --- /dev/null +++ b/src/core/auth/mod.rs @@ -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), + } + } +} diff --git a/src/core/databases/conversations.rs b/src/core/databases/conversations.rs new file mode 100644 index 0000000..7e5ffb6 --- /dev/null +++ b/src/core/databases/conversations.rs @@ -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, + pub updated_at: chrono::DateTime, +} + +#[derive(Debug)] +pub struct ConversationList { + pub conversations: Vec, + pub has_more: bool, +} + +pub struct CursorPage { + pub limit: u32, + pub before: Option>, +} + +#[derive(Debug)] +pub struct MessageSummary { + pub id: Uuid, + pub parent_id: Option, + pub role: ChatRole, + pub content: String, + pub created_at: chrono::DateTime, + pub tokens: Option, +} + +#[derive(Debug)] +pub struct MessageList { + pub messages: Vec, + pub has_more: bool, +} + +#[derive(Debug)] +pub enum ConversationResult { + Existing(Uuid), + Created(Uuid), +} diff --git a/src/core/databases/mod.rs b/src/core/databases/mod.rs new file mode 100644 index 0000000..7fa3605 --- /dev/null +++ b/src/core/databases/mod.rs @@ -0,0 +1 @@ +pub mod conversations; diff --git a/src/core/llm/chat.rs b/src/core/llm/chat.rs new file mode 100644 index 0000000..f4e37ab --- /dev/null +++ b/src/core/llm/chat.rs @@ -0,0 +1,79 @@ +use super::ChatRole; + +use crate::providers::ollama::errors::LlmError; + +use futures::Stream; +use serde::Serialize; +use std::pin::Pin; + +#[derive(Debug, Clone, Default)] +pub struct ChatCompletionOptions { + pub seed: Option, + pub temperature: Option, + pub top_p: Option, + pub top_k: Option, + pub stop: Option>, + pub num_ctx: Option, + pub num_predict: Option, + + pub keep_alive: Option, + 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, + pub parent_id: Option, +} + +#[derive(Debug, Clone, Serialize)] +pub struct ChatCompletionResultNoStream { + pub model: String, + pub message: Message, + + pub created_at: String, + pub id: uuid::Uuid, + + pub prompt_tokens: u32, + pub completion_tokens: u32, + + pub done_reason: Option, + + pub total_duration: Option, + pub load_duration: Option, +} + +#[derive(Serialize)] +pub enum ChatCompletionStreamEvent { + Token(String), + Final(ChatCompletionResultNoStream), +} + +pub type ChatCompletionStream = + Pin> + 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, + } + } +} diff --git a/src/core/llm/completions.rs b/src/core/llm/completions.rs new file mode 100644 index 0000000..20abe1f --- /dev/null +++ b/src/core/llm/completions.rs @@ -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, + pub temperature: Option, + pub top_p: Option, + pub top_k: Option, + pub stop: Option>, + pub num_ctx: Option, + pub num_predict: Option, + + pub keep_alive: Option, + 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, + + pub total_duration: Option, + pub load_duration: Option, +} + +#[derive(Serialize)] +pub enum CompletionStreamEvent { + Token(String), + Final(CompletionResultNoStream), +} + +pub type CompletionStream = + Pin> + Send>>; + +pub enum CompletionResult { + Stream(CompletionStream), + NoStream(CompletionResultNoStream), +} diff --git a/src/core/llm/mod.rs b/src/core/llm/mod.rs new file mode 100644 index 0000000..9405660 --- /dev/null +++ b/src/core/llm/mod.rs @@ -0,0 +1,6 @@ +pub mod chat; +pub mod completions; +pub mod models; +pub mod role; + +pub use role::ChatRole; diff --git a/src/core/llm/models.rs b/src/core/llm/models.rs new file mode 100644 index 0000000..805d7a7 --- /dev/null +++ b/src/core/llm/models.rs @@ -0,0 +1,36 @@ +use std::collections::HashMap; + +#[derive(Debug, Clone)] +pub struct Models { + pub models: Vec, +} + +#[derive(Debug, Clone)] +pub struct Model { + pub name: String, + + /// Optional grouping (OpenAI = "gpt", Ollama = "llama", etc.) + pub family: Option, + + /// Human-readable size like "7B", "13B", "gpt-4" + pub size: Option, + + /// Optional metadata, provider-specific info normalized into a string map + pub metadata: ModelMetadata, +} + +#[derive(Debug, Clone, Default)] +pub struct ModelMetadata { + pub extra: HashMap, +} + +#[derive(Debug, Clone)] +pub struct LoadModelRequest { + pub model: String, + pub keep_alive: String, +} + +#[derive(Debug, Clone)] +pub struct LoadModelResponse { + pub model: String, +} diff --git a/src/core/llm/role.rs b/src/core/llm/role.rs new file mode 100644 index 0000000..ca03fb5 --- /dev/null +++ b/src/core/llm/role.rs @@ -0,0 +1,8 @@ +use serde::Serialize; + +#[derive(Debug, Clone, Serialize)] +pub enum ChatRole { + System, + User, + Assistant, +} diff --git a/src/core/mod.rs b/src/core/mod.rs new file mode 100644 index 0000000..b119627 --- /dev/null +++ b/src/core/mod.rs @@ -0,0 +1,3 @@ +pub mod auth; +pub mod databases; +pub mod llm; diff --git a/src/databases/postgres/errors.rs b/src/databases/errors.rs similarity index 82% rename from src/databases/postgres/errors.rs rename to src/databases/errors.rs index 6534304..de37aa4 100644 --- a/src/databases/postgres/errors.rs +++ b/src/databases/errors.rs @@ -8,6 +8,9 @@ pub enum DbError { #[error("database timeout")] Timeout, + #[error("not authorized")] + Unauthorized, + #[error("not found")] NotFound, } diff --git a/src/databases/mod.rs b/src/databases/mod.rs index 26e9103..a069935 100644 --- a/src/databases/mod.rs +++ b/src/databases/mod.rs @@ -1 +1,3 @@ pub mod postgres; + +pub mod errors; diff --git a/src/databases/postgres/api_key/mod.rs b/src/databases/postgres/api_key/mod.rs index 84c032e..0333ab5 100644 --- a/src/databases/postgres/api_key/mod.rs +++ b/src/databases/postgres/api_key/mod.rs @@ -1 +1,2 @@ pub mod queries; +pub mod types; diff --git a/src/databases/postgres/api_key/queries.rs b/src/databases/postgres/api_key/queries.rs index d38b7d7..90c6e62 100644 --- a/src/databases/postgres/api_key/queries.rs +++ b/src/databases/postgres/api_key/queries.rs @@ -1,9 +1,11 @@ -use crate::databases::postgres::errors::DbError; +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> { +pub async fn update_last_access(pool: &PgPool, api_key_id: &Uuid) -> Result<(), DbError> { sqlx::query!( r#" UPDATE auth.api_key @@ -17,3 +19,41 @@ pub async fn update_last_access(pool: &PgPool, api_key_id: Uuid) -> Result<(), D 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 { + 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" + 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) +} diff --git a/src/databases/postgres/api_key/types.rs b/src/databases/postgres/api_key/types.rs new file mode 100644 index 0000000..cc7801e --- /dev/null +++ b/src/databases/postgres/api_key/types.rs @@ -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, +} + +#[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, +} diff --git a/src/databases/postgres/chat/queries.rs b/src/databases/postgres/chat/queries.rs index facc050..a845d58 100644 --- a/src/databases/postgres/chat/queries.rs +++ b/src/databases/postgres/chat/queries.rs @@ -1,10 +1,10 @@ -use crate::databases::postgres::chat::types; -use crate::databases::postgres::errors::DbError; +use crate::databases::errors::DbError; +use crate::databases::postgres::chat::{types, types::MessageRole}; use sqlx::{Acquire, PgPool}; use uuid::Uuid; -// ---- Creation ---- +// ---- Helpers ---- async fn validate_conversation<'e, E>(executor: E, conversation_id: Uuid) -> Result<(), DbError> where @@ -40,14 +40,14 @@ where Ok(rec.id) } +// ------ Creation ------ + pub async fn get_or_create_conversation( pool: &PgPool, conversation_id: Option, user_id: Uuid, ) -> Result { - tracing::debug!("Testing conversation"); - - let mut tx = pool.begin().await?; + let mut tx: sqlx::Transaction<'_, sqlx::Postgres> = pool.begin().await?; let result = { let conn = tx.acquire().await?; @@ -165,7 +165,7 @@ pub async fn update_message_tokens( pub async fn get_conversations_entries( pool: &PgPool, user_id: Uuid, - limit: i64, + limit: u32, before: Option>, ) -> Result, DbError> { let mut tx = pool.begin().await?; @@ -187,7 +187,7 @@ pub async fn get_conversations_entries( "#, user_id, before, - limit + limit as i64 ) .fetch_all(&mut *tx) .await?; @@ -200,7 +200,7 @@ pub async fn get_conversation_messages( pool: &PgPool, user_id: Uuid, conversation_id: Uuid, - limit: i64, + limit: u32, before: Option>, ) -> Result, DbError> { let mut conn = pool.acquire().await?; @@ -213,16 +213,16 @@ pub async fn get_conversation_messages( let rows = sqlx::query_as!( types::MessageSummary, r#" - SELECT id, parent_id, role, content, created_at, tokens + 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 ASC + ORDER BY created_at DESC LIMIT $3 "#, conversation_id, before, - limit + limit as i64 ) .fetch_all(&mut *conn) .await?; diff --git a/src/databases/postgres/chat/types.rs b/src/databases/postgres/chat/types.rs index c2f0583..93c957e 100644 --- a/src/databases/postgres/chat/types.rs +++ b/src/databases/postgres/chat/types.rs @@ -1,8 +1,7 @@ -use serde::Serialize; use uuid::Uuid; #[derive(Debug, Clone, sqlx::Type)] -#[sqlx(type_name = "text")] +#[sqlx(type_name = "chat.role")] #[sqlx(rename_all = "lowercase")] pub enum MessageRole { User, @@ -16,7 +15,7 @@ pub enum ConversationState { Created(Uuid), } -#[derive(Debug, sqlx::FromRow, Serialize)] +#[derive(Debug, sqlx::FromRow)] pub struct ConversationSummary { pub id: Uuid, pub title: Option, @@ -24,11 +23,11 @@ pub struct ConversationSummary { pub updated_at: chrono::DateTime, } -#[derive(Debug, sqlx::FromRow, Serialize)] +#[derive(Debug, sqlx::FromRow)] pub struct MessageSummary { pub id: Uuid, pub parent_id: Option, - pub role: String, + pub role: MessageRole, pub content: String, pub created_at: chrono::DateTime, pub tokens: Option, diff --git a/src/databases/postgres/mod.rs b/src/databases/postgres/mod.rs index 4db98f0..8d7898d 100644 --- a/src/databases/postgres/mod.rs +++ b/src/databases/postgres/mod.rs @@ -1,6 +1,5 @@ -pub mod errors; pub mod pool; pub mod api_key; pub mod chat; -pub mod user; +pub mod user_activity; diff --git a/src/databases/postgres/pool.rs b/src/databases/postgres/pool.rs index 9608cba..3ef3e20 100644 --- a/src/databases/postgres/pool.rs +++ b/src/databases/postgres/pool.rs @@ -1,4 +1,4 @@ -use super::errors::DbError; +use crate::databases::errors::DbError; use sqlx::{PgPool, postgres::PgPoolOptions}; use std::time::Duration; diff --git a/src/databases/postgres/user/queries.rs b/src/databases/postgres/user/queries.rs deleted file mode 100644 index 8e6a0ef..0000000 --- a/src/databases/postgres/user/queries.rs +++ /dev/null @@ -1,20 +0,0 @@ -use crate::databases::postgres::errors::DbError; - -use sqlx::PgPool; -use uuid::Uuid; - -pub async fn ensure_user_exists(pool: &PgPool, user_id: Uuid) -> Result<(), DbError> { - sqlx::query!( - r#" - INSERT INTO auth.app_user (id) - VALUES ($1) - ON CONFLICT (id) DO NOTHING - "#, - user_id - ) - .execute(pool) - .await - .map_err(DbError::Connection)?; - - Ok(()) -} diff --git a/src/databases/postgres/user/mod.rs b/src/databases/postgres/user_activity/mod.rs similarity index 100% rename from src/databases/postgres/user/mod.rs rename to src/databases/postgres/user_activity/mod.rs diff --git a/src/databases/postgres/user_activity/queries.rs b/src/databases/postgres/user_activity/queries.rs new file mode 100644 index 0000000..9fc4ce2 --- /dev/null +++ b/src/databases/postgres/user_activity/queries.rs @@ -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(()) +} diff --git a/src/docs.rs b/src/docs.rs deleted file mode 100644 index bd4d32a..0000000 --- a/src/docs.rs +++ /dev/null @@ -1,55 +0,0 @@ -use utoipa::OpenApi; - -use crate::dto::api; -use crate::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::chat::completions, - routes::v1::chat::chat_completions, - routes::v1::models::list_models, - routes::v1::models::load_model, - routes::v1::models::unload_model, - ), - components( - schemas( - api::ErrorResponse, - api::ModelsResponse, - api::ModelInfo, - api::LoadModelResponse, - api::LoadModelBody, - api::UnloadModelResponse, - api::BaseLLMRequest, - api::CompletionRequest, - api::CompletionObject, - api::FinishReason, - api::CompletionResponse, - api::Choice, - api::Usage, - api::CompletionChunk, - api::ChatRequest, - api::Message, - api::Role, - api::ChatCompletionResponse, - api::ChatChoice, - api::ChatCompletionChunk, - api::ChatChunkChoice, - api::Delta, - ) - ), - tags( - (name = "chat", description = "Chat & completions"), - (name = "models", description = "Model management") - ) -)] -pub struct ApiDoc; diff --git a/src/dto/mod.rs b/src/dto/mod.rs deleted file mode 100644 index b6fe8b2..0000000 --- a/src/dto/mod.rs +++ /dev/null @@ -1,2 +0,0 @@ -pub mod api; -pub mod ollama; diff --git a/src/dto/ollama.rs b/src/dto/ollama.rs deleted file mode 100644 index 6625748..0000000 --- a/src/dto/ollama.rs +++ /dev/null @@ -1,80 +0,0 @@ -use serde::{Deserialize, Serialize}; - -use crate::dto::api; - -#[derive(Debug, Serialize, Deserialize)] -pub struct OllamaModels { - pub models: Vec, -} - -#[derive(Debug, Serialize, Deserialize)] -pub struct OllamaModel { - pub name: String, - - pub details: Option, - - pub size: Option, - pub digest: Option, - pub modified_at: Option, -} - -#[derive(Debug, Serialize, Deserialize)] -pub struct OllamaModelDetails { - pub family: Option, - pub parameter_size: Option, - pub quantization_level: Option, -} - -#[derive(Debug, Serialize, Deserialize)] -pub struct OllamaOptions { - pub temperature: Option, - pub top_p: Option, - pub top_k: Option, - pub repeat_penalty: Option, - pub seed: Option, - - pub num_ctx: Option, - pub num_predict: Option, -} - -#[derive(Debug, Serialize)] -pub struct OllamaGenerateRequest<'a> { - pub model: &'a str, - pub prompt: &'a str, - pub stream: bool, - pub options: OllamaOptions, -} - -#[derive(Debug, Serialize, Deserialize)] -pub struct OllamaGenerateResponse { - pub model: String, - pub created_at: Option, - pub response: String, - pub done: bool, - - #[serde(default)] - pub context: Option>, - - pub total_duration: Option, - pub load_duration: Option, - pub prompt_eval_count: Option, - pub eval_count: Option, -} - -#[derive(Debug, Serialize)] -pub struct OllamaChatRequest<'a> { - pub model: &'a str, - pub messages: &'a [api::Message], - pub stream: bool, - pub options: OllamaOptions, -} - -#[derive(Debug, Deserialize)] -pub struct OllamaChatResponse { - pub model: String, - pub message: api::Message, - pub done: bool, - - pub prompt_eval_count: Option, - pub eval_count: Option, -} diff --git a/src/lib.rs b/src/lib.rs index ea8d223..c2ecda0 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,7 +1,6 @@ pub mod api; +pub mod core; pub mod databases; -pub mod dto; -pub mod middlewares; +pub mod mappers; pub mod providers; -pub mod state; -pub mod utils; +pub mod services; diff --git a/src/main.rs b/src/main.rs index b5f7acb..bc2429c 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,79 +1,21 @@ mod api; +mod core; mod databases; -mod docs; -mod dto; -mod middlewares; +mod mappers; mod providers; -mod routes; -mod state; -mod utils; +mod services; -use crate::databases::postgres; -use crate::providers::ollama::client::OllamaProvider; -use crate::state::app_state::AppState; - -use axum::Router; -use axum::http::{HeaderName, HeaderValue, Method, header}; -use once_cell::sync::Lazy; -use std::env; +use api::app::build_app; use std::net::SocketAddr; -use std::sync::Arc; -use tower_http::cors::CorsLayer; -use tracing_subscriber::{EnvFilter, fmt}; - -static OLLAMA_URL: Lazy = Lazy::new(|| env::var("OLLAMA_URL").expect("OLLAMA_URL not set")); - -pub fn init_tracing() { - let filter = env::var("RUST_LOG").unwrap_or_else(|_| "info".to_string()); - - fmt().with_env_filter(EnvFilter::new(filter)).init(); -} #[tokio::main] async fn main() { #[cfg(debug_assertions)] - { - dotenvy::dotenv().ok(); - } + dotenvy::dotenv().ok(); - init_tracing(); + tracing_subscriber::fmt::init(); - // DB Connection - 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 state = AppState { - ollama: Arc::new(OllamaProvider::new(OLLAMA_URL.as_str())), - postgres: pool, - }; - - let cors_origin = - env::var("CORS_ORIGIN").unwrap_or_else(|_| "http://localhost:3000".to_string()); - - let cors = CorsLayer::new() - .allow_origin(cors_origin.parse::().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); - - let app = Router::new() - .nest("/v1", routes::v1::router(state.clone())) - .layer(cors) - .with_state(state); + let app = build_app().await; let addr = SocketAddr::from(([0, 0, 0, 0], 3001)); tracing::debug!("Server running on {}", addr); diff --git a/src/mappers/api_to_core.rs b/src/mappers/api_to_core.rs new file mode 100644 index 0000000..b589f7d --- /dev/null +++ b/src/mappers/api_to_core.rs @@ -0,0 +1,66 @@ +use crate::{api, core}; + +impl From for core::llm::completions::CompletionRequest { + fn from(m: api::types::CompletionRequest) -> 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 for core::llm::chat::ChatCompletionRequest { + fn from(m: api::types::ChatRequest) -> 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 for core::llm::chat::Message { + fn from(m: api::types::Message) -> Self { + Self { + role: match m.role { + api::types::Role::System => core::llm::ChatRole::System, + api::types::Role::User => core::llm::ChatRole::User, + api::types::Role::Assistant => core::llm::ChatRole::Assistant, + }, + content: m.content, + } + } +} + +impl From 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, + } + } +} diff --git a/src/mappers/core_to_api.rs b/src/mappers/core_to_api.rs new file mode 100644 index 0000000..2886675 --- /dev/null +++ b/src/mappers/core_to_api.rs @@ -0,0 +1,101 @@ +use crate::{api, core}; + +impl From for api::types::ModelsResponse { + fn from(m: core::llm::models::Models) -> Self { + Self { + models: m.models.into_iter().map(Into::into).collect(), + } + } +} + +impl From for api::types::ModelMetadata { + fn from(m: core::llm::models::ModelMetadata) -> Self { + Self { + extra: m + .extra + .iter() + .map(|(k, v)| (k.clone(), v.clone())) + .collect(), + } + } +} + +impl From for api::types::ModelInfo { + 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 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 + 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 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 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 for api::types::CompletionResponse { + fn from(m: core::llm::completions::CompletionResultNoStream) -> Self { + Self { + id: m.id, + object: api::types::CompletionObject::TextCompletion, + created: m.created_at, + model: m.model, + choices: vec![api::types::Choice { + finish_reason: api::types::FinishReason::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, + }, + } + } +} diff --git a/src/mappers/core_to_database.rs b/src/mappers/core_to_database.rs new file mode 100644 index 0000000..9c87062 --- /dev/null +++ b/src/mappers/core_to_database.rs @@ -0,0 +1,30 @@ +impl From 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 + 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 + } + } + } +} diff --git a/src/mappers/core_to_ollama.rs b/src/mappers/core_to_ollama.rs new file mode 100644 index 0000000..c476041 --- /dev/null +++ b/src/mappers/core_to_ollama.rs @@ -0,0 +1,67 @@ +use crate::core; +use crate::providers::ollama; + +impl From 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 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 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 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 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()), + } + } +} diff --git a/src/mappers/database_to_core.rs b/src/mappers/database_to_core.rs new file mode 100644 index 0000000..abaabe0 --- /dev/null +++ b/src/mappers/database_to_core.rs @@ -0,0 +1,69 @@ +use crate::core; +use crate::databases::postgres; + +impl From + 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 + 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 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 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 + 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) + } + } + } +} diff --git a/src/mappers/keycloak_to_core.rs b/src/mappers/keycloak_to_core.rs new file mode 100644 index 0000000..7489577 --- /dev/null +++ b/src/mappers/keycloak_to_core.rs @@ -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 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::>(); + + 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, + } + } +} diff --git a/src/mappers/mod.rs b/src/mappers/mod.rs new file mode 100644 index 0000000..f818e67 --- /dev/null +++ b/src/mappers/mod.rs @@ -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; diff --git a/src/mappers/ollama_to_core.rs b/src/mappers/ollama_to_core.rs new file mode 100644 index 0000000..0a2eb83 --- /dev/null +++ b/src/mappers/ollama_to_core.rs @@ -0,0 +1,50 @@ +use crate::core; +use crate::providers::ollama; + +impl From for core::llm::models::Models { + fn from(m: ollama::types::OllamaModels) -> Self { + Self { + models: m.models.into_iter().map(Into::into).collect(), + } + } +} + +impl From 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 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, + } + } +} diff --git a/src/middlewares/auth/apikey.rs b/src/middlewares/auth/apikey.rs deleted file mode 100644 index 6304d10..0000000 --- a/src/middlewares/auth/apikey.rs +++ /dev/null @@ -1,15 +0,0 @@ -use uuid::Uuid; - -#[derive(PartialEq, Eq, Clone, Debug)] -pub enum ApiKeyClaimsRoles { - Admin, - Read, - Write, -} - -#[derive(Clone, Debug)] -pub struct ApiKeyClaims { - pub sub: Uuid, - pub api_key_id: Uuid, - pub roles: Vec, -} diff --git a/src/middlewares/auth/keycloak.rs b/src/middlewares/auth/keycloak.rs deleted file mode 100644 index 157f297..0000000 --- a/src/middlewares/auth/keycloak.rs +++ /dev/null @@ -1,148 +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::collections::HashMap; -use std::env; -use std::sync::Arc; -use std::time::{Duration, Instant}; -use tokio::sync::RwLock; - -// ------ JWKS ------ - -#[derive(Clone)] -struct JwksCache { - jwks: Value, - last_fetched: Instant, -} - -static JWK_CACHE: Lazy>>> = Lazy::new(|| Arc::new(RwLock::new(None))); - -static JWKS_URL: Lazy = Lazy::new(|| env::var("JWKS_URL").expect("JWKS_URL not set")); - -async fn fetch_jwks() -> Result { - let jwks = reqwest::get(JWKS_URL.as_str()) - .await? - .json::() - .await?; - - Ok(jwks) -} - -pub async fn refresh_jwks() -> Result { - let jwks = fetch_jwks().await?; - - let mut write = JWK_CACHE.write().await; - - *write = Some(JwksCache { - jwks: jwks.clone(), - last_fetched: Instant::now(), - }); - - Ok(jwks) -} - -pub async fn get_jwks() -> Result { - let ttl = Duration::from_secs(3600); // 1 hour - - { - // Read lock first (fast path) - let read = JWK_CACHE.read().await; - - if let Some(cache) = read - .as_ref() - .filter(|cache| cache.last_fetched.elapsed() < ttl) - { - return Ok(cache.jwks.clone()); - } - } - - // Expired or empty → refresh - refresh_jwks().await -} - -// ------ Claims ------ - -#[derive(Debug, Deserialize, Serialize, Clone, Default)] -pub struct KeycloakClaims { - pub sub: String, - pub preferred_username: Option, - pub exp: usize, - pub iss: String, - pub aud: Option>, - pub realm_access: Option, - #[serde(default)] - pub resource_access: HashMap, -} - -#[derive(Debug, Deserialize, Serialize, Clone)] -pub struct RealmAccess { - pub roles: Vec, -} - -#[derive(Debug, Deserialize, Serialize, Clone)] -pub struct ResourceAccess { - pub roles: Vec, -} - -impl KeycloakClaims { - pub fn realm_roles(&self) -> &[String] { - self.realm_access - .as_ref() - .map_or(&[], |r| r.roles.as_slice()) - } - - pub fn has_realm_role(&self, role: &str) -> bool { - self.realm_roles().iter().any(|r| r == role) - } - - pub fn client_roles(&self, client: &str) -> &[String] { - self.resource_access - .get(client) - .map_or(&[], |r| r.roles.as_slice()) - } - - pub fn has_client_role(&self, client: &str, role: &str) -> bool { - self.client_roles(client).iter().any(|r| r == role) - } -} - -// ------ Validation ------ - -static ISSUER: Lazy = Lazy::new(|| env::var("ISSUER").expect("ISSUER not set")); - -pub fn validate_token(token: &str, jwks: &Value) -> Result { - // 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()]); - - validation.validate_exp = true; - validation.validate_aud = false; - - // 5. Decode & verify - let token_data = decode::(token, &decoding_key, &validation) - .map_err(|_| "Token validation failed")?; - - Ok(token_data.claims) -} diff --git a/src/middlewares/auth/middleware.rs b/src/middlewares/auth/middleware.rs deleted file mode 100644 index a437cc4..0000000 --- a/src/middlewares/auth/middleware.rs +++ /dev/null @@ -1,227 +0,0 @@ -use axum::{ - extract::{Request, State}, - http::StatusCode, - middleware::Next, - response::{IntoResponse, Response}, -}; - -use crate::databases::postgres::{ - api_key::queries::update_last_access, user::queries::ensure_user_exists, -}; -use crate::middlewares::auth::apikey::{ApiKeyClaims, ApiKeyClaimsRoles}; -use crate::middlewares::auth::keycloak::{KeycloakClaims, get_jwks, refresh_jwks, validate_token}; -use crate::state::app_state::AppState; -use crate::utils::crypto::hash_key; -use uuid::Uuid; - -#[derive(Clone, Debug)] -pub enum Auth { - Jwt(KeycloakClaims), - ApiKey(ApiKeyClaims), -} - -impl Auth { - pub fn user_id(&self) -> Uuid { - match self { - Auth::Jwt(c) => c.sub.parse().expect("sub is a valid UUID"), - Auth::ApiKey(c) => c.sub, - } - } - - pub fn has_realm_role(&self, role: &str) -> bool { - match self { - Auth::Jwt(c) => c.has_realm_role(role), - Auth::ApiKey(_) => false, // API keys carry no roles - } - } - - pub fn has_client_role(&self, client: &str, role: &str) -> bool { - match self { - Auth::Jwt(c) => c.has_client_role(client, role), - Auth::ApiKey(_) => false, - } - } -} - -pub async fn auth_middleware( - State(state): State, - request: Request, - next: Next, -) -> Result { - match try_jwt(&state, request, next).await { - Ok(response) => Ok(response), - Err((request, next)) => try_api_key(&state, request, next).await, - } -} - -/// Returns Ok(Response) if JWT was valid and request handled. -/// Returns Err((request, next)) if no JWT was present (caller should try next method). -/// Returns a 401/500 response directly if JWT was present but invalid. -async fn try_jwt( - state: &AppState, - request: Request, - next: Next, -) -> Result { - let token = request - .headers() - .get("authorization") - .and_then(|v| v.to_str().ok()) - .and_then(|v| v.strip_prefix("Bearer ")) - .map(str::to_owned); - - let Some(token) = token else { - // No Authorization header at all → let API key branch try - return Err((request, next)); - }; - - let jwks = match get_jwks().await { - Ok(j) => j, - Err(e) => { - tracing::error!("Failed to fetch JWKS: {e}"); - // Token was present but we can't validate → hard 500 - return Ok(StatusCode::INTERNAL_SERVER_ERROR.into_response()); - } - }; - - let claims = match validate_token(&token, &jwks) { - Ok(c) => c, - Err(_) => { - // Try refreshing JWKS once - match refresh_jwks().await { - Ok(fresh_jwks) => match validate_token(&token, &fresh_jwks) { - Ok(c) => c, - Err(_) => { - tracing::warn!("JWT validation failed after JWKS refresh"); - return Ok(StatusCode::UNAUTHORIZED.into_response()); - } - }, - Err(e) => { - tracing::error!("Failed to refresh JWKS: {e}"); - return Ok(StatusCode::INTERNAL_SERVER_ERROR.into_response()); - } - } - } - }; - - tracing::debug!("JWT valid, sub={}", claims.sub); - handle_auth(state, request, next, Auth::Jwt(claims)) - .await - .map_err(|_| unreachable!()) -} - -/// Returns Ok(Response) if API key was valid. -/// Returns Err(StatusCode) otherwise (UNAUTHORIZED or INTERNAL_SERVER_ERROR). -async fn try_api_key( - state: &AppState, - request: Request, - next: Next, -) -> Result { - let key = request - .headers() - .get("x-api-key") - .and_then(|v| v.to_str().ok()) - .ok_or(StatusCode::UNAUTHORIZED)? - .to_owned(); - - let key_hash = hash_key(&key); - - let row = sqlx::query!( - r#" - SELECT u.id AS user_id, ak.id as key_id, ak.scopes as roles - 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(&state.postgres) - .await - .map_err(|e| { - tracing::error!("DB error during API key lookup: {e}"); - StatusCode::INTERNAL_SERVER_ERROR - })? - .ok_or(StatusCode::UNAUTHORIZED)?; - - tracing::debug!("API key valid, user_id={}", row.user_id); - - let roles = row - .roles - .into_iter() - .filter_map(|r| match r.as_str() { - "admin" => Some(ApiKeyClaimsRoles::Admin), - "read" => Some(ApiKeyClaimsRoles::Read), - "write" => Some(ApiKeyClaimsRoles::Write), - _ => None, - }) - .collect(); - - handle_auth( - state, - request, - next, - Auth::ApiKey(ApiKeyClaims { - sub: row.user_id, - api_key_id: row.key_id, - roles, - }), - ) - .await -} - -// ── Shared post-auth logic ──────────────────────────────────────────────────── - -/// Ensures the user exists in the DB, inserts `Auth` into extensions, runs the next handler. -async fn handle_auth( - state: &AppState, - mut request: Request, - next: Next, - auth: Auth, -) -> Result { - match &auth { - Auth::Jwt(_) => { - ensure_user_exists(&state.postgres, auth.user_id()) - .await - .map_err(|e| { - tracing::error!("ensure_user_exists failed: {e}"); - StatusCode::INTERNAL_SERVER_ERROR - })?; - } - Auth::ApiKey(key) => { - update_last_access(&state.postgres, key.api_key_id) - .await - .map_err(|e| { - tracing::error!("update_last_access failed: {e}"); - StatusCode::INTERNAL_SERVER_ERROR - })?; - } - } - - request.extensions_mut().insert(auth); - Ok(next.run(request).await) -} - -// ── Role guard ─────────────────────────────────────────────────────────────── - -/// Layer-level middleware that checks roles *after* `auth_middleware` has run. -pub async fn require_roles( - request: Request, - next: Next, - realm_role: Option<&'static str>, - client_role: Option<&'static str>, -) -> Result { - let auth = request - .extensions() - .get::() - .ok_or(StatusCode::UNAUTHORIZED)?; - - if realm_role.is_some_and(|role| !auth.has_realm_role(role)) { - return Err(StatusCode::FORBIDDEN); - } - - if client_role.is_some_and(|role| !auth.has_client_role("chat-api", role)) { - return Err(StatusCode::FORBIDDEN); - } - - Ok(next.run(request).await) -} diff --git a/src/middlewares/auth/mod.rs b/src/middlewares/auth/mod.rs deleted file mode 100644 index a8b3d11..0000000 --- a/src/middlewares/auth/mod.rs +++ /dev/null @@ -1,5 +0,0 @@ -pub mod apikey; -pub mod keycloak; -pub mod middleware; - -pub use middleware::auth_middleware; diff --git a/src/providers/keycloak/claims.rs b/src/providers/keycloak/claims.rs new file mode 100644 index 0000000..3158471 --- /dev/null +++ b/src/providers/keycloak/claims.rs @@ -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, + pub exp: usize, + pub iss: String, + pub _aud: Option>, + pub realm_access: Option, + pub resource_access: HashMap, +} + +#[derive(Debug, Clone, Deserialize)] +pub struct RealmAccess { + pub roles: Vec, +} + +#[derive(Debug, Clone, Deserialize)] +pub struct ResourceAccess { + pub roles: Vec, +} diff --git a/src/providers/keycloak/errors.rs b/src/providers/keycloak/errors.rs new file mode 100644 index 0000000..ba3bb2f --- /dev/null +++ b/src/providers/keycloak/errors.rs @@ -0,0 +1,40 @@ +use thiserror::Error; + +#[derive(Debug, Error)] +pub enum AuthError { + #[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), +} diff --git a/src/providers/keycloak/jwks.rs b/src/providers/keycloak/jwks.rs new file mode 100644 index 0000000..144803e --- /dev/null +++ b/src/providers/keycloak/jwks.rs @@ -0,0 +1,59 @@ +use super::errors::AuthError; + +use once_cell::sync::Lazy; +use serde_json::Value; +use std::env; +use std::sync::Arc; +use std::time::{Duration, Instant}; +use tokio::sync::RwLock; + +#[derive(Clone)] +struct JwksCache { + jwks: Value, + last_fetched: Instant, +} + +static JWK_CACHE: Lazy>>> = Lazy::new(|| Arc::new(RwLock::new(None))); + +static JWKS_URL: Lazy = Lazy::new(|| env::var("JWKS_URL").expect("JWKS_URL not set")); + +async fn fetch_jwks() -> Result { + let jwks = reqwest::get(JWKS_URL.as_str()) + .await? + .json::() + .await?; + + Ok(jwks) +} + +pub async fn refresh_jwks() -> Result { + let jwks = fetch_jwks().await?; + + let mut write = JWK_CACHE.write().await; + + *write = Some(JwksCache { + jwks: jwks.clone(), + last_fetched: Instant::now(), + }); + + Ok(jwks) +} + +pub async fn get_jwks() -> Result { + let ttl = Duration::from_secs(3600); // 1 hour + + { + // Read lock first (fast path) + let read = JWK_CACHE.read().await; + + if let Some(cache) = read + .as_ref() + .filter(|cache| cache.last_fetched.elapsed() < ttl) + { + return Ok(cache.jwks.clone()); + } + } + + // Expired or empty → refresh + refresh_jwks().await +} diff --git a/src/providers/keycloak/mod.rs b/src/providers/keycloak/mod.rs new file mode 100644 index 0000000..864dee4 --- /dev/null +++ b/src/providers/keycloak/mod.rs @@ -0,0 +1,4 @@ +pub mod claims; +pub mod errors; +pub mod jwks; +pub mod validator; diff --git a/src/providers/keycloak/validator.rs b/src/providers/keycloak/validator.rs new file mode 100644 index 0000000..295bb9f --- /dev/null +++ b/src/providers/keycloak/validator.rs @@ -0,0 +1,64 @@ +use super::claims::KeycloakClaims; +use super::errors::AuthError; + +use jsonwebtoken::{Algorithm, DecodingKey, Validation, decode, decode_header}; +use once_cell::sync::Lazy; +use serde_json::Value; +use std::env; + +static ISSUER: Lazy = Lazy::new(|| env::var("ISSUER").expect("ISSUER not set")); + +fn validate_token(token: &str, jwks: &Value) -> Result { + // 1. Decode header + let header = decode_header(token).map_err(|_| AuthError::InvalidHeader)?; + + let kid = header.kid.ok_or(AuthError::MissingKid)?; + + // 2. Find matching key + let keys = jwks["keys"].as_array().ok_or(AuthError::InvalidJwks)?; + + let key = keys + .iter() + .find(|k| k["kid"] == kid) + .ok_or(AuthError::InvalidDecodingKey)?; + + // 3. Extract RSA components + let n = key["n"].as_str().ok_or(AuthError::MissingModulus)?; + let e = key["e"].as_str().ok_or(AuthError::MissingExponent)?; + + let decoding_key = + DecodingKey::from_rsa_components(n, e).map_err(|_| AuthError::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::(token, &decoding_key, &validation) + .map_err(|_| AuthError::TokenValidationFailed)?; + + Ok(token_data.claims) +} + +pub async fn authenticate_jwt(token: &str) -> Result { + let jwks = super::jwks::get_jwks() + .await + .map_err(|_| AuthError::JwksFetchFailed)?; + + match validate_token(token, &jwks) { + Ok(claims) => Ok(claims), + + Err(_) => { + // one retry with refresh + let fresh = super::jwks::refresh_jwks() + .await + .map_err(|_| AuthError::JwksRefreshFailed)?; + + validate_token(token, &fresh).map_err(|_| AuthError::InvalidToken) + } + } +} diff --git a/src/providers/mod.rs b/src/providers/mod.rs index eb9349e..c7df8a0 100644 --- a/src/providers/mod.rs +++ b/src/providers/mod.rs @@ -1 +1,2 @@ +pub mod keycloak; pub mod ollama; diff --git a/src/providers/ollama/client.rs b/src/providers/ollama/client.rs index 99a7273..95f9995 100644 --- a/src/providers/ollama/client.rs +++ b/src/providers/ollama/client.rs @@ -1,10 +1,8 @@ -use super::errors::OllamaError; -use crate::dto::{api, ollama}; -use axum::response::sse::Event; +use crate::providers::ollama; +use crate::providers::ollama::errors::LlmError; + use futures::StreamExt; use reqwest::Client; -use serde_json::json; -use tokio_stream::wrappers::ReceiverStream; #[derive(Clone)] pub struct OllamaProvider { @@ -22,7 +20,7 @@ impl OllamaProvider { // ── private helpers ────────────────────────────────────────────────────── - async fn model_exists(&self, model: &str) -> Result { + async fn model_exists(&self, model: &str) -> Result { let url = format!("{}/api/tags", self.base_url); let res = self @@ -30,61 +28,43 @@ impl OllamaProvider { .get(url) .send() .await? - .json::() + .json::() .await?; Ok(res.models.iter().any(|m| m.name == model)) } - fn has_user_message(&self, messages: &[api::Message]) -> bool { - messages.iter().any(|m| matches!(m.role, api::Role::User)) + fn has_user_message(&self, messages: &[ollama::types::OllamaMessage]) -> bool { + messages + .iter() + .any(|m| matches!(m.role, ollama::types::OllamaRole::User)) } - fn extract_completion_params<'a>( - &self, - body: &'a api::CompletionRequest, - ) -> Result<(&'a str, &'a str), OllamaError> { - let prompt = body.prompt.trim(); + // pub fn validate_keep_alive(&self, s: &str) -> Result<(), LlmError> { + // let s = s.trim(); - let model = body.base.model.as_str(); + // if s == "-1" || s.parse::().is_ok() { + // return Ok(()); + // } - Ok((prompt, model)) - } + // let split = s + // .find(|c: char| c.is_alphabetic()) + // .ok_or_else(|| LlmError::InvalidKeepAlive(s.to_string()))?; - fn extract_chat_params<'a>( - &self, - body: &'a api::ChatRequest, - ) -> Result<(&'a [api::Message], &'a str), OllamaError> { - let model = body.base.model.as_str(); + // let (num, unit) = s.split_at(split); - Ok((&body.messages, model)) - } + // num.parse::() + // .map_err(|_| LlmError::InvalidKeepAlive(s.to_string()))?; - pub fn parse_keep_alive(&self, s: &str) -> Result<(), OllamaError> { - let s = s.trim(); - - if s == "-1" || s.parse::().is_ok() { - return Ok(()); - } - - let split = s - .find(|c: char| c.is_alphabetic()) - .ok_or_else(|| OllamaError::InvalidKeepAlive(s.to_string()))?; - - let (num, unit) = s.split_at(split); - - num.parse::() - .map_err(|_| OllamaError::InvalidKeepAlive(s.to_string()))?; - - match unit { - "s" | "m" | "h" => Ok(()), - _ => Err(OllamaError::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 { + pub async fn list_models(&self) -> Result { let url = format!("{}/api/tags", self.base_url); let res = self @@ -92,156 +72,83 @@ impl OllamaProvider { .get(url) .send() .await? - .json::() + .json::() .await?; - let models = res.models.into_iter().map(api::ModelInfo::from).collect(); - - Ok(api::ModelsResponse { models }) - } - - pub async fn load_model( - &self, - model: &str, - keep_alive: Option<&str>, - ) -> Result { - let url = format!("{}/api/generate", self.base_url); - - let keep_alive = keep_alive.ok_or(OllamaError::MissingKeepAlive)?; - - self.parse_keep_alive(keep_alive)?; - - let exists = self.model_exists(model).await?; - if !exists { - return Err(OllamaError::ModelNotFound(model.to_string())); - } - - let payload = json!({ - "model": model, - "prompt": "", - "keep_alive": keep_alive, - "stream": false, - }); - - let _res = self - .client - .post(url) - .json(&payload) - .send() - .await? - .text() - .await?; - - Ok(api::LoadModelResponse { - model: model.to_string(), - status: "loaded".to_string(), - keep_alive: keep_alive.to_string(), - }) - } - - pub async fn unload_model(&self, model: &str) -> Result { - let url = format!("{}/api/generate", self.base_url); - - let exists = self.model_exists(model).await?; - if !exists { - return Err(OllamaError::ModelNotFound(model.to_string())); - } - - let payload = json!({ - "model": model, - "prompt": "", - "keep_alive": "0", - "stream": false, - }); - - let _res = self - .client - .post(url) - .json(&payload) - .send() - .await? - .text() - .await?; - - Ok(api::UnloadModelResponse { - model: model.to_string(), - status: "unloaded".to_string(), - }) + Ok(res) } pub async fn completions( &self, - body: &api::CompletionRequest, - ) -> Result { + body: &super::types::OllamaGenerateRequest, + ) -> Result { let url = format!("{}/api/generate", self.base_url); - let (prompt, model) = self.extract_completion_params(body)?; - - if prompt.is_empty() { - return Err(OllamaError::MissingPrompt); + if body.prompt.is_empty() { + return Err(LlmError::MissingPrompt); } - if model.is_empty() { - return Err(OllamaError::MissingModel); + if body.model.is_empty() { + return Err(LlmError::MissingModel); } - let exists = self.model_exists(model).await?; + let exists = self.model_exists(&body.model).await?; if !exists { - return Err(OllamaError::ModelNotFound(model.to_string())); + return Err(LlmError::ModelNotFound(body.model.clone())); } - let options = ollama::OllamaOptions::from(&body.base); + let options = body.options.clone(); - let payload = ollama::OllamaGenerateRequest { - model, - prompt, + 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::() + .json::() .await?; - Ok(api::CompletionResponse::from(res)) + Ok(res) } pub async fn completions_stream( &self, - body: &api::CompletionRequest, - ) -> Result>, OllamaError> { + body: &super::types::OllamaGenerateRequest, + ) -> Result { let url = format!("{}/api/generate", self.base_url); - let (prompt, model) = self.extract_completion_params(body)?; - - if prompt.is_empty() { - return Err(OllamaError::MissingPrompt); + if body.prompt.is_empty() { + return Err(LlmError::MissingPrompt); } - if model.is_empty() { - return Err(OllamaError::MissingModel); + if body.model.is_empty() { + return Err(LlmError::MissingModel); } - let exists = self.model_exists(model).await?; + let exists = self.model_exists(&body.model.to_string()).await?; if !exists { - return Err(OllamaError::ModelNotFound(model.to_string())); + return Err(LlmError::ModelNotFound(body.model.to_string())); } - let options = ollama::OllamaOptions::from(&body.base); - - let payload = ollama::OllamaGenerateRequest { - model, - prompt, + let payload = ollama::types::OllamaGenerateRequest { + model: body.model.clone(), + prompt: body.prompt.clone(), stream: true, - options, + keep_alive: body.keep_alive.clone(), + options: body.options.clone(), }; - let mut byte_stream = self + let byte_stream = self .client .post(url) .json(&payload) @@ -249,85 +156,80 @@ impl OllamaProvider { .await? .bytes_stream(); - let (tx, rx) = tokio::sync::mpsc::channel(32); + let stream = byte_stream.flat_map(|chunk_result| { + let mut out: Vec> = + Vec::new(); - tokio::spawn(async move { - while let Some(chunk) = byte_stream.next().await { - let chunk = match chunk { - Ok(b) => b, - Err(e) => { - let _ = tx.send(Err(OllamaError::Http(e))).await; - break; - } - }; + let chunk = match chunk_result { + Ok(b) => b, + Err(e) => { + out.push(Err(LlmError::Http(e))); + return futures::stream::iter(out); + } + }; - // 🔥 IMPORTANT: typed deserialization - let parsed: ollama::OllamaGenerateResponse = match serde_json::from_slice(&chunk) { - Ok(v) => v, - Err(_) => continue, - }; + for line in chunk.split(|&b| b == b'\n') { + if line.is_empty() { + continue; + } - // map → OpenAI chunk - let event_data = serde_json::to_string(&api::CompletionChunk { - id: "cmpl-ollama".to_string(), - object: "text_completion".to_string(), - choices: vec![api::Choice { - text: parsed.response, - index: 0, - finish_reason: if parsed.done { - api::FinishReason::Stop - } else { - api::FinishReason::Length - }, - }], - }) - .unwrap_or_default(); + let parsed: ollama::types::OllamaGenerateResponse = + match serde_json::from_slice(line) { + Ok(v) => v, + Err(_) => continue, + }; - let _ = tx.send(Ok(Event::default().data(event_data))).await; + if !parsed.response.is_empty() && !parsed.done { + out.push(Ok(super::types::OllamaGenerateStreamEvent::Token( + parsed.response.clone(), + ))); + } if parsed.done { - let _ = tx.send(Ok(Event::default().data("[DONE]"))).await; - break; + out.push(Ok(super::types::OllamaGenerateStreamEvent::Final(parsed))); + return futures::stream::iter(out); } } + + futures::stream::iter(out) }); - Ok(ReceiverStream::new(rx)) + Ok(Box::pin(stream)) } pub async fn chat_completions( &self, - body: &api::ChatRequest, - ) -> Result { + body: &super::types::OllamaChatRequest, + ) -> Result { let url = format!("{}/api/chat", self.base_url); - let (messages, model) = self.extract_chat_params(body)?; - if body.messages.is_empty() { - return Err(OllamaError::MissingMessages); + return Err(LlmError::MissingMessages); } - if !self.has_user_message(&body.messages) { - return Err(OllamaError::MissingMessages); + let ollama_messages: Vec = body.messages.clone(); + + if !self.has_user_message(&ollama_messages) { + return Err(LlmError::MissingMessages); } - if model.is_empty() { - return Err(OllamaError::MissingModel); + if body.model.is_empty() { + return Err(LlmError::MissingModel); } - let exists = self.model_exists(model).await?; + let exists = self.model_exists(&body.model).await?; if !exists { - return Err(OllamaError::ModelNotFound(model.to_string())); + return Err(LlmError::ModelNotFound(body.model.clone())); } - // let options = ollama::OllamaOptions::from(body); - let options = ollama::OllamaOptions::from(&body.base); + let options = body.options.clone(); - let payload = ollama::OllamaChatRequest { - model, - messages, + let payload = ollama::types::OllamaChatRequest { + model: body.model.clone(), + messages: ollama_messages, stream: false, options, + keep_alive: body.keep_alive.clone(), }; let res = self @@ -336,43 +238,46 @@ impl OllamaProvider { .json(&payload) .send() .await? - .json::() + .json::() .await?; - Ok(api::ChatCompletionResponse::from(res)) + Ok(res) } pub async fn chat_completions_stream( &self, - body: &api::ChatRequest, - ) -> Result>, OllamaError> { + body: &super::types::OllamaChatRequest, + ) -> Result { let url = format!("{}/api/chat", self.base_url); - let (messages, model) = self.extract_chat_params(body)?; - - if messages.is_empty() { - return Err(OllamaError::MissingMessages); + if body.messages.is_empty() { + return Err(LlmError::MissingMessages); } - if model.is_empty() { - return Err(OllamaError::MissingModel); + if body.model.is_empty() { + return Err(LlmError::MissingModel); } - let exists = self.model_exists(model).await?; + let exists = self.model_exists(&body.model).await?; if !exists { - return Err(OllamaError::ModelNotFound(model.to_string())); + return Err(LlmError::ModelNotFound(body.model.clone())); } - let options = ollama::OllamaOptions::from(&body.base); + let ollama_messages: Vec = body.messages.clone(); - let payload = ollama::OllamaChatRequest { - model, - messages, + 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, + options: body.options.clone(), + keep_alive: body.keep_alive.clone(), }; - let mut byte_stream = self + let byte_stream = self .client .post(url) .json(&payload) @@ -380,63 +285,42 @@ impl OllamaProvider { .await? .bytes_stream(); - let (tx, rx) = tokio::sync::mpsc::channel(32); + let stream = byte_stream.flat_map(|chunk_result| { + let mut out: Vec> = Vec::new(); - tokio::spawn(async move { - let stream_id = format!("chatcmpl-{}", uuid::Uuid::new_v4()); + let chunk = match chunk_result { + Ok(b) => b, + Err(e) => { + out.push(Err(LlmError::Http(e))); + return futures::stream::iter(out); + } + }; - while let Some(chunk) = byte_stream.next().await { - let chunk = match chunk { - Ok(b) => b, - Err(e) => { - let _ = tx.send(Err(OllamaError::Http(e))).await; - break; - } - }; + for line in chunk.split(|&b| b == b'\n') { + if line.is_empty() { + continue; + } - let parsed: ollama::OllamaChatResponse = match serde_json::from_slice(&chunk) { + let parsed: ollama::types::OllamaChatResponse = match serde_json::from_slice(line) { Ok(v) => v, Err(_) => continue, }; - let usage = if parsed.done { - Some(api::Usage { - prompt_tokens: parsed.prompt_eval_count.unwrap_or(0) as u32, - completion_tokens: parsed.eval_count.unwrap_or(0) as u32, - total_tokens: (parsed.prompt_eval_count.unwrap_or(0) - + parsed.eval_count.unwrap_or(0)) - as u32, - }) - } else { - None - }; - - let event = api::ChatCompletionChunk { - id: stream_id.clone(), - object: "chat.completion.chunk".to_string(), - choices: vec![api::ChatChunkChoice { - index: 0, - delta: api::Delta { - role: Some(parsed.message.role), - content: Some(parsed.message.content), - }, - finish_reason: if parsed.done { - Some(api::FinishReason::Stop) - } else { - None - }, - }], - usage, - }; - - let _ = tx.send(Ok(event)).await; + if !parsed.message.content.is_empty() && !parsed.done { + out.push(Ok(super::types::OllamaChatStreamEvent::Token( + parsed.message.content.clone(), + ))); + } if parsed.done { - break; + out.push(Ok(super::types::OllamaChatStreamEvent::Final(parsed))); + return futures::stream::iter(out); } } + + futures::stream::iter(out) }); - Ok(ReceiverStream::new(rx)) + Ok(Box::pin(stream)) } } diff --git a/src/providers/ollama/errors.rs b/src/providers/ollama/errors.rs index 2484a95..a3ce800 100644 --- a/src/providers/ollama/errors.rs +++ b/src/providers/ollama/errors.rs @@ -1,8 +1,7 @@ -// errors.rs use thiserror::Error; #[derive(Debug, Error)] -pub enum OllamaError { +pub enum LlmError { #[error("prompt is required and cannot be empty")] MissingPrompt, @@ -15,44 +14,10 @@ pub enum OllamaError { #[error("model '{0}' is not available — run `ollama pull {0}` first")] ModelNotFound(String), - #[error("keep_alive is required and cannot be empty")] - MissingKeepAlive, - - #[error( - "invalid keep_alive format '{0}' — expected (e.g. 30s, 10m, 2h), a plain integer (seconds), or -1" - )] - InvalidKeepAlive(String), - + // #[error( + // "invalid keep_alive format '{0}' — expected (e.g. 30s, 10m, 2h), a plain integer (seconds), or -1" + // )] + // InvalidKeepAlive(String), #[error(transparent)] Http(#[from] reqwest::Error), } - -pub fn into_http_response(e: OllamaError) -> (axum::http::StatusCode, String) { - match e { - OllamaError::MissingPrompt => ( - axum::http::StatusCode::BAD_REQUEST, - "prompt is required and cannot be empty".to_string(), - ), - OllamaError::MissingModel => ( - axum::http::StatusCode::BAD_REQUEST, - "model is required and cannot be empty".to_string(), - ), - OllamaError::ModelNotFound(m) => ( - axum::http::StatusCode::UNPROCESSABLE_ENTITY, - format!("model '{m}' is not available — run `ollama pull {m}` first"), - ), - OllamaError::MissingKeepAlive => ( - axum::http::StatusCode::BAD_REQUEST, - "keep alive is required and cannot be empty".to_string(), - ), - OllamaError::InvalidKeepAlive(v) => ( - axum::http::StatusCode::BAD_REQUEST, - format!("invalid keep_alive '{v}'"), - ), - OllamaError::MissingMessages => ( - axum::http::StatusCode::BAD_REQUEST, - "messages array with at least one user message is required".to_string(), - ), - OllamaError::Http(e) => (axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string()), - } -} diff --git a/src/providers/ollama/mapper.rs b/src/providers/ollama/mapper.rs deleted file mode 100644 index 18a4ae0..0000000 --- a/src/providers/ollama/mapper.rs +++ /dev/null @@ -1,88 +0,0 @@ -use crate::dto::{api, ollama}; - -use chrono::Utc; -use uuid::Uuid; - -impl From for api::ModelInfo { - fn from(m: ollama::OllamaModel) -> Self { - Self { - name: m.name, - - family: m.details.as_ref().and_then(|d| d.family.clone()), - parameter_size: m.details.as_ref().and_then(|d| d.parameter_size.clone()), - quantization: m - .details - .as_ref() - .and_then(|d| d.quantization_level.clone()), - } - } -} - -impl From for api::CompletionResponse { - fn from(res: ollama::OllamaGenerateResponse) -> Self { - Self { - id: Uuid::new_v4().to_string(), - object: api::CompletionObject::TextCompletion, - model: res.model, - created: Utc::now().timestamp() as u64, - - choices: vec![api::Choice { - text: res.response, - index: 0, - finish_reason: api::FinishReason::Stop, - }], - - usage: api::Usage { - prompt_tokens: res.prompt_eval_count.unwrap_or(0), - completion_tokens: res.eval_count.unwrap_or(0), - total_tokens: res.prompt_eval_count.unwrap_or(0) + res.eval_count.unwrap_or(0), - }, - } - } -} - -impl From<&api::BaseLLMRequest> for ollama::OllamaOptions { - fn from(base: &api::BaseLLMRequest) -> Self { - Self { - temperature: base.temperature, - top_p: base.top_p, - top_k: base.top_k, - repeat_penalty: base.repeat_penalty, - seed: base.seed, - num_ctx: base.num_ctx, - num_predict: base.num_predict, - } - } -} - -impl From for api::ChatCompletionResponse { - fn from(res: ollama::OllamaChatResponse) -> Self { - let prompt_tokens = res.prompt_eval_count.unwrap_or(0); - let completion_tokens = res.eval_count.unwrap_or(0); - - Self { - id: Uuid::new_v4().to_string(), - object: "chat.completion".to_string(), - created: Utc::now().timestamp() as u64, - model: res.model, - - choices: vec![api::ChatChoice { - index: 0, - message: res.message, - finish_reason: if res.done { - api::FinishReason::Stop - } else { - api::FinishReason::Length - }, - }], - - usage: Some(api::Usage { - prompt_tokens, - completion_tokens, - total_tokens: prompt_tokens + completion_tokens, - }), - - conversation_id: None, - } - } -} diff --git a/src/providers/ollama/mod.rs b/src/providers/ollama/mod.rs index d04f432..5963cd0 100644 --- a/src/providers/ollama/mod.rs +++ b/src/providers/ollama/mod.rs @@ -1,3 +1,3 @@ pub mod client; pub mod errors; -pub mod mapper; +pub mod types; diff --git a/src/providers/ollama/types.rs b/src/providers/ollama/types.rs new file mode 100644 index 0000000..b991917 --- /dev/null +++ b/src/providers/ollama/types.rs @@ -0,0 +1,137 @@ +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, +} + +#[derive(Debug, Deserialize)] +pub struct OllamaModel { + pub name: String, + + pub details: Option, + + pub size: Option, + pub digest: Option, + pub modified_at: Option, +} + +#[derive(Debug, Deserialize)] +pub struct OllamaModelDetails { + pub family: Option, + pub parameter_size: Option, + pub quantization_level: Option, +} + +// ------ 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, + pub temperature: Option, + pub top_p: Option, + pub top_k: Option, + pub stop: Option>, + pub num_ctx: Option, + pub num_predict: Option, +} + +// ------ Completion ------ + +#[derive(Debug, Serialize)] +pub struct OllamaGenerateRequest { + pub model: String, + pub prompt: String, + pub stream: bool, + pub keep_alive: String, + pub options: Option, +} + +#[derive(Debug, Deserialize)] +pub struct OllamaGenerateResponse { + pub model: String, + pub created_at: String, + + pub response: String, + + pub done: bool, + pub done_reason: Option, + + pub total_duration: Option, + pub load_duration: Option, + + pub prompt_eval_count: Option, + pub eval_count: Option, +} + +#[derive(Debug)] +pub enum OllamaGenerateStreamEvent { + Token(String), + Final(OllamaGenerateResponse), +} + +pub type OllamaGenerateResponseStream = Pin< + Box< + dyn Stream> + Send, + >, +>; + +// ------ Chat ------ + +#[derive(Debug, Serialize)] +pub struct OllamaChatRequest { + pub model: String, + pub messages: Vec, + pub stream: bool, + pub keep_alive: String, + pub options: Option, +} + +#[derive(Debug, Deserialize)] +pub struct OllamaChatResponse { + pub model: String, + pub created_at: String, + + pub message: OllamaMessage, + + pub done: bool, + pub done_reason: Option, + + pub total_duration: Option, + pub load_duration: Option, + + pub prompt_eval_count: Option, + pub eval_count: Option, +} + +#[derive(Debug)] +pub enum OllamaChatStreamEvent { + Token(String), + Final(OllamaChatResponse), +} + +pub type OllamaChatResponseStream = Pin< + Box> + Send>, +>; diff --git a/src/routes/v1/apikey.rs b/src/routes/v1/apikey.rs deleted file mode 100644 index 4835ddd..0000000 --- a/src/routes/v1/apikey.rs +++ /dev/null @@ -1,50 +0,0 @@ -use axum::{ - Json, - extract::{Extension, State}, - http::StatusCode, -}; -use base64::{Engine as _, engine::general_purpose}; -use rand::RngCore; -use rand::rngs::OsRng; - -use crate::dto::api::{CreateApiKeyRequest, CreateApiKeyResponse}; -use crate::middlewares::auth::apikey::ApiKeyClaimsRoles; -use crate::middlewares::auth::middleware::Auth; -use crate::state::app_state::AppState; -use crate::utils::crypto::hash_key; - -fn generate_api_key() -> String { - let mut bytes = [0u8; 32]; - OsRng.fill_bytes(&mut bytes); - general_purpose::URL_SAFE_NO_PAD.encode(bytes) -} - -pub async fn create_api_key( - State(state): State, - Extension(claims): Extension, - Json(body): Json, -) -> Result, StatusCode> { - if matches!(&claims, Auth::ApiKey(api_key) if !api_key.roles.contains(&ApiKeyClaimsRoles::Admin)) - { - return Err(StatusCode::FORBIDDEN); - } - - let raw_key = generate_api_key(); - let key_hash = hash_key(&raw_key); - - sqlx::query!( - r#" - INSERT INTO auth.api_key (key_hash, name, created_by, scopes) - VALUES ($1, $2, $3, $4) - "#, - key_hash, - body.name, - claims.user_id(), - &body.scopes - ) - .execute(&state.postgres) - .await - .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; - - Ok(Json(CreateApiKeyResponse { api_key: raw_key })) -} diff --git a/src/routes/v1/chat.rs b/src/routes/v1/chat.rs deleted file mode 100644 index 83491bb..0000000 --- a/src/routes/v1/chat.rs +++ /dev/null @@ -1,504 +0,0 @@ -use crate::{ - dto::api::BaseLLMRequest, middlewares::auth::middleware::Auth, - providers::ollama::client::OllamaProvider, -}; -use axum::{ - Json, - extract::{Extension, Path, Query, State}, - response::{ - IntoResponse, Response, - sse::{Event, KeepAlive, Sse}, - }, -}; -use sqlx::PgPool; -use tokio_stream::StreamExt; -use uuid::Uuid; - -use crate::api::errors::ApiError; -use crate::databases::postgres::chat::queries::{ - get_conversation_messages, get_conversations_entries, get_or_create_conversation, - insert_message, set_conversation_title, update_message_tokens, -}; -use crate::databases::postgres::chat::types::{ConversationState, MessageRole}; -use crate::databases::postgres::errors; -use crate::dto::api; -use crate::providers::ollama::errors::into_http_response; -use crate::state::app_state::AppState; - -#[utoipa::path( - post, - path = "/completions", - tag = "chat", - request_body( - content = api::CompletionRequest, - 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::CompletionResponse, - content_type = "application/json" - ), - ( - status = 400, - description = "Invalid request: missing prompt, model, or invalid format", - body = api::ErrorResponse, - example = json!({ "error": "prompt is required and cannot be empty" }) - ), - ( - status = 422, - description = "Model not found or not available locally", - body = api::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 = api::ErrorResponse, - example = json!({ "error": "connection refused" }) - ) - ) -)] -pub async fn completions( - State(state): State, - Json(body): Json, -) -> Result)> { - tracing::debug!("Received /completion with body {:?}", body); - if body.base.stream { - let stream = state.ollama.completions_stream(&body).await.map_err(|e| { - let (code, msg) = into_http_response(e); - (code, Json(api::ErrorResponse::new(msg))) - })?; - - Ok(Sse::new(stream) - .keep_alive(KeepAlive::default()) - .into_response()) - } else { - let response = state.ollama.completions(&body).await.map_err(|e| { - let (code, msg) = into_http_response(e); - (code, Json(api::ErrorResponse::new(msg))) - })?; - - Ok(Json(response).into_response()) - } -} - -async fn ensure_conversation( - state: &AppState, - auth: &Auth, - conversation_id: Option, - first_message: &str, -) -> Result { - let conversation_state = - get_or_create_conversation(&state.postgres, conversation_id, auth.user_id()).await?; - - let id = match conversation_state { - ConversationState::Existing(uuid) => uuid, - ConversationState::Created(uuid) => { - let pool = state.postgres.clone(); - let ollama = state.ollama.clone(); - let user_id = auth.user_id(); - let first_message = first_message.to_string(); - - tokio::spawn(async move { - let title = generate_conversation_title(&ollama, &first_message).await; - if let Err(e) = set_conversation_title(&pool, uuid, user_id, &title).await { - tracing::warn!( - conversation_id = %uuid, - error = %e, - "Failed to set conversation title" - ); - } - tracing::debug!(conversation_id = %uuid, %title, "generated conversation title"); - }); - - uuid - } - }; - - Ok(id) -} - -async fn handle_stream( - state: AppState, - auth: Auth, - body: api::ChatRequest, - conversation_id: Option, -) -> Result { - // Handle anonymous (API key) path early — no DB logging - let Some(conv_id) = conversation_id else { - let stream = state.ollama.chat_completions_stream(&body).await?; - - let plain_stream = stream.map( - |item| -> Result { - match item { - Ok(chunk) => Ok( - Event::default().data(serde_json::to_string(&chunk).unwrap_or_default()) - ), - Err(e) => Err(e), - } - }, - ); - return Ok(Sse::new(plain_stream) - .keep_alive(KeepAlive::default()) - .into_response()); - }; - - // From here conv_id is a plain Uuid — all variables stay in scope - let user_msg_id = log_user_message( - &state.postgres, - auth.user_id(), - conv_id, - body.parent_id, - body.messages - .last() - .map(|m| m.content.as_str()) - .unwrap_or(""), - None, - ) - .await?; - - let start_event = api::StreamEvent::Start(api::StartEventData { - conversation_id: conv_id, - created: chrono::Utc::now().timestamp() as u64, - id: user_msg_id, - }); - - let (tx, rx) = tokio::sync::mpsc::channel::< - Result, - >(32); - - // Send start event immediately, before Ollama is contacted - let _ = tx - .send(Ok(Event::default() - .event("metadata") - .data(serde_json::to_string(&start_event).unwrap()))) - .await; - - let pool = state.postgres.clone(); - let user_id = auth.user_id(); - - tokio::spawn(async move { - // Ollama called inside spawn — start event already queued - let stream = match state.ollama.chat_completions_stream(&body).await { - Ok(s) => s, - Err(e) => { - let _ = tx.send(Err(e)).await; - return; - } - }; - - let mut stream = stream; - let mut accumulated = String::new(); - - while let Some(item) = futures::StreamExt::next(&mut stream).await { - match item { - Ok(chunk) => { - let is_done = chunk.choices[0].finish_reason == Some(api::FinishReason::Stop); - - if let Some(content) = chunk.choices[0].delta.content.as_ref() { - accumulated.push_str(content); - } - - if is_done { - let prompt_tokens = chunk.usage.as_ref().map(|u| u.prompt_tokens); - let completion_tokens = chunk.usage.as_ref().map(|u| u.completion_tokens); - - if let Some(pt) = prompt_tokens { - let _ = update_message_tokens(&pool, user_id, user_msg_id, pt).await; - } - - let assistant_msg_id = log_assistant_message( - &pool, - user_id, - conv_id, - user_msg_id, - &accumulated, - completion_tokens, - ) - .await; - - if let Ok(msg_id) = assistant_msg_id { - let end_event = api::StreamEvent::End(api::EndEventData { - usage: api::Usage { - prompt_tokens: prompt_tokens.unwrap_or(0), - completion_tokens: completion_tokens.unwrap_or(0), - total_tokens: chunk - .usage - .as_ref() - .map(|u| u.total_tokens) - .unwrap_or(0), - }, - id: msg_id, - created: chrono::Utc::now().timestamp() as u64, - }); - - let _ = tx - .send(Ok(Event::default() - .event("metadata") - .data(serde_json::to_string(&end_event).unwrap()))) - .await; - } - - break; - } - - let data = api::StreamEvent::Delta(chunk); - let json = serde_json::to_string(&data).unwrap(); - if tx.send(Ok(Event::default().data(json))).await.is_err() { - break; - } - } - Err(e) => { - let _ = tx.send(Err(e)).await; - break; - } - } - } - }); - - Ok(Sse::new(tokio_stream::wrappers::ReceiverStream::new(rx)) - .keep_alive(KeepAlive::default()) - .into_response()) -} - -async fn handle_non_stream( - state: AppState, - auth: Auth, - body: api::ChatRequest, - conversation_id: Option, -) -> Result { - let mut response = state.ollama.chat_completions(&body).await?; - - response.conversation_id = conversation_id; - - if let Some(conversation_id) = conversation_id { - let user_msg_id = log_user_message( - &state.postgres, - auth.user_id(), - conversation_id, - body.parent_id, - body.messages - .last() - .map(|m| m.content.as_str()) - .unwrap_or(""), - response.usage.map(|u| u.prompt_tokens), - ) - .await?; - - log_assistant_message( - &state.postgres, - auth.user_id(), - conversation_id, - user_msg_id, - &response.choices[0].message.content, - response.usage.map(|u| u.completion_tokens), - ) - .await?; - } - - Ok(Json(response).into_response()) -} - -#[utoipa::path( - post, - path = "/chat/completions", - tag = "chat", - request_body( - content = api::ChatRequest, - 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::ChatCompletionResponse, - content_type = "application/json" - ), - ( - status = 400, - description = "Invalid request", - body = api::ErrorResponse, - example = json!({ "error": "messages array with at least one user message is required" }) - ), - ( - status = 401, - description = "Unauthorized", - body = api::ErrorResponse, - example = json!({ "error": "missing or invalid token" }) - ), - ( - status = 422, - description = "Model not found or unavailable", - body = api::ErrorResponse, - example = json!({ "error": "model 'llama3' is not available — run `ollama pull llama3` first" }) - ), - ( - status = 500, - description = "Internal server error", - body = api::ErrorResponse, - example = json!({ "error": "connection refused" }) - ) - ) -)] -pub async fn chat_completions( - State(state): State, - Extension(auth): Extension, - Json(mut body): Json, -) -> Result { - tracing::debug!("Received /chat/completion with body {:?}", body); - - let conversation_id = if matches!(&auth, Auth::Jwt(_)) { - let first_message = body.messages[0].content.clone(); - let conv_id = - ensure_conversation(&state, &auth, body.conversation_id, &first_message).await?; - - if let Some(depth) = body.base.context_depth { - dbg!({ depth }); - if depth > 0 { - let history = get_conversation_messages( - &state.postgres, - auth.user_id(), - conv_id, - depth as i64, - None, // no cursor — fetch the most recent N messages - ) - .await?; - - // Map MessageSummary → api::Message and prepend to the outgoing request - let history_messages: Vec = history - .into_iter() - .map(|m| api::Message { - role: api::Role::Assistant, // TODO - content: m.content, - }) - .collect(); - - // body.messages = [history_messages, body.messages].concat(); - body.messages.splice(0..0, history_messages); - } - } - - dbg!("{:?}", &body.messages); - - Some(conv_id) - } else { - None - }; - - if body.base.stream { - handle_stream(state, auth, body, conversation_id).await - } else { - handle_non_stream(state, auth, body, conversation_id).await - } -} - -async fn generate_conversation_title(ollama: &OllamaProvider, first_message: &str) -> String { - let request = api::CompletionRequest { - base: BaseLLMRequest { - model: "llama3:latest".to_string(), - ..Default::default() - }, - prompt: format!( - "Generate a short title (max 6 words) for the following chat conversation: {}", - first_message - ), - }; - - ollama - .completions(&request) - .await - .ok() - .and_then(|r| r.choices.first().map(|c| c.text.trim().to_string())) - .filter(|t| !t.is_empty()) - .unwrap_or_else(|| "New Conversation".to_string()) // ← default on any error -} - -async fn log_user_message( - pool: &PgPool, - user_id: Uuid, - conversation_id: Uuid, - parent_id: Option, - content: &str, - tokens: Option, -) -> Result { - insert_message( - pool, - user_id, - conversation_id, - parent_id, - MessageRole::User, - content, - tokens, - ) - .await -} - -async fn log_assistant_message( - pool: &PgPool, - user_id: Uuid, - conversation_id: Uuid, - parent_id: Uuid, - content: &str, - tokens: Option, -) -> Result { - insert_message( - pool, - user_id, - conversation_id, - Some(parent_id), - MessageRole::Assistant, - content, - tokens, - ) - .await -} - -// Conversation retrieveing -pub async fn get_conversations( - State(state): State, - Extension(auth): Extension, - Query(params): Query, -) -> Result, ApiError> { - tracing::debug!("Conversation hit: {:?}", auth); - - let conversations = get_conversations_entries( - &state.postgres, - auth.user_id(), - params.limit.unwrap_or(20), - params.before, - ) - .await?; - - let has_more = conversations.len() == params.limit.unwrap_or(20) as usize; - - Ok(Json(api::ConversationListResponse { - conversations, - has_more, - })) -} - -pub async fn get_messages( - State(state): State, - Extension(auth): Extension, - Path(conversation_id): Path, - Query(params): Query, -) -> Result, ApiError> { - tracing::debug!("Messages hit: {:?}", auth); - - let messages = get_conversation_messages( - &state.postgres, - auth.user_id(), - conversation_id, - params.limit.unwrap_or(50), - params.before, - ) - .await?; - - let has_more = messages.len() == params.limit.unwrap_or(50) as usize; - - Ok(Json(api::MessageListResponse { messages, has_more })) -} diff --git a/src/routes/v1/models.rs b/src/routes/v1/models.rs deleted file mode 100644 index 39ae278..0000000 --- a/src/routes/v1/models.rs +++ /dev/null @@ -1,134 +0,0 @@ -use axum::{ - Json, - extract::{Path, State}, -}; - -use crate::dto::api; -use crate::providers::ollama::errors::into_http_response; -use crate::state::app_state::AppState; - -#[utoipa::path( - get, - path = "/models", - tag = "models", - responses( - ( - status = 200, - description = "List of locally available Ollama models", - body = api::ModelsResponse, - content_type = "application/json", - ), - ( - status = 500, - description = "Internal server error (Ollama or network failure)", - body = api::ErrorResponse, - example = json!({ "error": "connection refused" }) - ) - ) -)] -pub async fn list_models( - State(state): State, -) -> Result, (axum::http::StatusCode, String)> { - match state.ollama.list_models().await { - Ok(models) => Ok(Json(models)), - Err(e) => Err(into_http_response(e)), - } -} - -#[utoipa::path( - post, - path = "/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::LoadModelBody, - description = "Load model request", - content_type = "application/json", - example = json!({ "keep_alive": "10m" }) - ), - responses( - ( - status = 200, - description = "Model successfully loaded into memory", - body = api::LoadModelResponse, - content_type = "application/json", - ), - ( - status = 400, - description = "Invalid or missing keep_alive format", - body = api::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::ErrorResponse, - example = json!({ "error": "model 'llama3' not found — run `ollama pull llama3`" }) - ), - ( - status = 500, - description = "Internal server error (Ollama or network failure)", - body = api::ErrorResponse, - example = json!({ "error": "connection refused" }) - ) - ) -)] -pub async fn load_model( - State(state): State, - Path(model): Path, - Json(body): Json, -) -> Result, (axum::http::StatusCode, String)> { - let response = state - .ollama - .load_model(&model, body.keep_alive.as_deref()) - .await - .map_err(into_http_response)?; - - Ok(Json(response)) -} - -#[utoipa::path( - delete, - path = "/models/{model}/load", - 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::UnloadModelResponse, - content_type = "application/json", - ), - ( - status = 404, - description = "Model not found locally", - body = api::ErrorResponse, - example = json!({ "error": "model 'llama3' not found — run `ollama pull llama3`" }) - ), - ( - status = 500, - description = "Internal server error (Ollama or network failure)", - body = api::ErrorResponse, - example = json!({ "error": "connection refused" }) - ) - ) -)] -pub async fn unload_model( - State(state): State, - Path(model): Path, -) -> Result, (axum::http::StatusCode, String)> { - let response = state - .ollama - .unload_model(&model) - .await - .map_err(into_http_response)?; - - Ok(Json(response)) -} diff --git a/src/routes/v1/openapi.rs b/src/routes/v1/openapi.rs deleted file mode 100644 index b5e5508..0000000 --- a/src/routes/v1/openapi.rs +++ /dev/null @@ -1,8 +0,0 @@ -// use axum::Json; -// use utoipa::OpenApi; - -// use crate::openapi::V1ApiDoc; - -// pub async fn openapi_json() -> Json { -// Json(V1ApiDoc::openapi()) -// } diff --git a/src/services/auth_service.rs b/src/services/auth_service.rs new file mode 100644 index 0000000..f9bacd7 --- /dev/null +++ b/src/services/auth_service.rs @@ -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 { + 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 { + 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 { + let auth = crate::providers::keycloak::validator::authenticate_jwt(token).await?; + + Ok(auth.into()) + } +} diff --git a/src/services/chat_service.rs b/src/services/chat_service.rs new file mode 100644 index 0000000..845fad8 --- /dev/null +++ b/src/services/chat_service.rs @@ -0,0 +1,317 @@ +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 futures::StreamExt; +use std::boxed::Box; +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 { + let models = self.ollama.list_models().await?; + + Ok(models.into()) + } + + pub async fn load_model( + &self, + body: crate::core::llm::models::LoadModelRequest, + ) -> Result { + 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 complete( + &self, + body: core::llm::completions::CompletionRequest, + ) -> Result { + 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 { + let conversation_id = self + .resolve_conversation_with_title(auth.user_id(), &body) + .await?; + + let user_msg_id = self + .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.message.clone(), + 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 ollama_stream = self.ollama.chat_completions_stream(&request).await?; + + let mapped = ollama_stream.map(|item| { + item.map(|event| match event { + OllamaChatStreamEvent::Token(tok) => { + core::llm::chat::ChatCompletionStreamEvent::Token(tok) + } + OllamaChatStreamEvent::Final(resp) => { + core::llm::chat::ChatCompletionStreamEvent::Final( + core::llm::chat::ChatCompletionResultNoStream { + id: Uuid::new_v4(), + created_at: resp.created_at, + model: resp.model, + message: resp.message.into(), + 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(core::llm::chat::ChatCompletionResult::Stream(Box::pin( + mapped, + ))) + } else { + let response = self.ollama.chat_completions(&request).await?; + + self.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.unwrap_or_default(), + ) + .await?; + + let enriched = core::llm::chat::ChatCompletionResultNoStream { + id: Uuid::new_v4(), + created_at: response.created_at, + model: response.model, + message: response.message.into(), + 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::chat::ChatCompletionResult::NoStream(enriched)) + } + } + + // ------ Helpers ------ + + async fn resolve_conversation_with_title( + &self, + user_id: Uuid, + body: &core::llm::chat::ChatCompletionRequest, + ) -> Result { + 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, + body_messages: crate::core::llm::chat::Message, + context_depth: u32, + ) -> Result, 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(); + + history.push(body_messages); + + Ok(history) + } + + async fn log_user_message( + &self, + user_id: Uuid, + conversation_id: Uuid, + parent_id: Option, + content: &str, + tokens: Option, + ) -> Result { + self.conversation + .log_message( + user_id, + conversation_id, + parent_id, + llm::ChatRole::User, + content, + tokens, + ) + .await + } + + async fn log_assistant_message( + &self, + user_id: Uuid, + conversation_id: Uuid, + parent_id: Uuid, + content: &str, + tokens: Option, + ) -> Result { + self.conversation + .log_message( + user_id, + conversation_id, + Some(parent_id), + llm::ChatRole::Assistant, + content, + tokens, + ) + .await + } +} diff --git a/src/services/conversation_service.rs b/src/services/conversation_service.rs new file mode 100644 index 0000000..5e378dd --- /dev/null +++ b/src/services/conversation_service.rs @@ -0,0 +1,132 @@ +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 { + 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 { + let limit = pointer.limit; + + let messages = postgres::chat::queries::get_conversation_messages( + &self.postgres, + user_id, + conversation_id, + limit, + pointer.before, + ) + .await?; + + let has_more = messages.len() == limit as usize; + + 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, + ) -> Result { + 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, + role: core::llm::ChatRole, + content: &str, + tokens: Option, + ) -> Result { + 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 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(()) + } +} diff --git a/src/services/errors.rs b/src/services/errors.rs new file mode 100644 index 0000000..0b61500 --- /dev/null +++ b/src/services/errors.rs @@ -0,0 +1,27 @@ +use crate::databases::errors::DbError; +use crate::providers::keycloak::errors::AuthError; +use crate::providers::ollama::errors::LlmError; + +pub enum ServiceError { + Db(DbError), + Llm(LlmError), + Auth(AuthError), +} + +impl From for ServiceError { + fn from(e: DbError) -> Self { + ServiceError::Db(e) + } +} + +impl From for ServiceError { + fn from(e: LlmError) -> Self { + ServiceError::Llm(e) + } +} + +impl From for ServiceError { + fn from(e: AuthError) -> Self { + ServiceError::Auth(e) + } +} diff --git a/src/services/mod.rs b/src/services/mod.rs new file mode 100644 index 0000000..5748bca --- /dev/null +++ b/src/services/mod.rs @@ -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; diff --git a/src/state/app_state.rs b/src/state/app_state.rs deleted file mode 100644 index 6a13b22..0000000 --- a/src/state/app_state.rs +++ /dev/null @@ -1,9 +0,0 @@ -use crate::providers::ollama::client::OllamaProvider; -use sqlx::PgPool; -use std::sync::Arc; - -#[derive(Clone)] -pub struct AppState { - pub ollama: Arc, - pub postgres: PgPool, -} diff --git a/src/state/mod.rs b/src/state/mod.rs deleted file mode 100644 index 384f36c..0000000 --- a/src/state/mod.rs +++ /dev/null @@ -1 +0,0 @@ -pub mod app_state; diff --git a/src/utils/crypto.rs b/src/utils/crypto.rs deleted file mode 100644 index c9c0df2..0000000 --- a/src/utils/crypto.rs +++ /dev/null @@ -1,11 +0,0 @@ -use sha2::{Digest, Sha256}; - -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() -} diff --git a/src/utils/mod.rs b/src/utils/mod.rs deleted file mode 100644 index 274f0ed..0000000 --- a/src/utils/mod.rs +++ /dev/null @@ -1 +0,0 @@ -pub mod crypto; diff --git a/tests/ollama_provider.rs b/tests/ollama_provider.rs index a5f3827..8878a4c 100644 --- a/tests/ollama_provider.rs +++ b/tests/ollama_provider.rs @@ -1,419 +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::>() - }) -} - -// ── 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(_))); -} +// 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::>() +// }) +// } + +// // ── 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(_))); +// }