diff --git a/README.md b/README.md index 6e2bffed..6d40e5a8 100644 --- a/README.md +++ b/README.md @@ -8,7 +8,7 @@ Honcho is a platform for making AI agents and LLM powered applications that are personalized to their end users. -Read about the motivation of this project [here](https://blog.plasticlabs.ai). +Read about the motivation of this project [here](https://blog.plasticlabs.ai/blog/A-Simple-Honcho-Primer). Read the user documenation [here](https://docs.honcho.dev) diff --git a/api/poetry.lock b/api/poetry.lock index 34d8d279..1bed420b 100644 --- a/api/poetry.lock +++ b/api/poetry.lock @@ -761,6 +761,55 @@ http2 = ["h2 (>=3,<5)"] socks = ["socksio (>=1.0.0,<2.0.0)"] trio = ["trio (>=0.22.0,<0.25.0)"] +[[package]] +name = "httptools" +version = "0.6.1" +description = "A collection of framework independent HTTP protocol utils." +category = "main" +optional = false +python-versions = ">=3.8.0" +files = [ + {file = "httptools-0.6.1-cp310-cp310-macosx_10_9_universal2.whl", hash = "sha256:d2f6c3c4cb1948d912538217838f6e9960bc4a521d7f9b323b3da579cd14532f"}, + {file = "httptools-0.6.1-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:00d5d4b68a717765b1fabfd9ca755bd12bf44105eeb806c03d1962acd9b8e563"}, + {file = "httptools-0.6.1-cp310-cp310-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:639dc4f381a870c9ec860ce5c45921db50205a37cc3334e756269736ff0aac58"}, + {file = "httptools-0.6.1-cp310-cp310-manylinux_2_5_x86_64.manylinux1_x86_64.manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:e57997ac7fb7ee43140cc03664de5f268813a481dff6245e0075925adc6aa185"}, + {file = "httptools-0.6.1-cp310-cp310-musllinux_1_1_aarch64.whl", hash = "sha256:0ac5a0ae3d9f4fe004318d64b8a854edd85ab76cffbf7ef5e32920faef62f142"}, + {file = "httptools-0.6.1-cp310-cp310-musllinux_1_1_x86_64.whl", hash = "sha256:3f30d3ce413088a98b9db71c60a6ada2001a08945cb42dd65a9a9fe228627658"}, + {file = "httptools-0.6.1-cp310-cp310-win_amd64.whl", hash = "sha256:1ed99a373e327f0107cb513b61820102ee4f3675656a37a50083eda05dc9541b"}, + {file = "httptools-0.6.1-cp311-cp311-macosx_10_9_universal2.whl", hash = "sha256:7a7ea483c1a4485c71cb5f38be9db078f8b0e8b4c4dc0210f531cdd2ddac1ef1"}, + {file = "httptools-0.6.1-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:85ed077c995e942b6f1b07583e4eb0a8d324d418954fc6af913d36db7c05a5a0"}, + {file = "httptools-0.6.1-cp311-cp311-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:8b0bb634338334385351a1600a73e558ce619af390c2b38386206ac6a27fecfc"}, + {file = "httptools-0.6.1-cp311-cp311-manylinux_2_5_x86_64.manylinux1_x86_64.manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:7d9ceb2c957320def533671fc9c715a80c47025139c8d1f3797477decbc6edd2"}, + {file = "httptools-0.6.1-cp311-cp311-musllinux_1_1_aarch64.whl", hash = "sha256:4f0f8271c0a4db459f9dc807acd0eadd4839934a4b9b892f6f160e94da309837"}, + {file = "httptools-0.6.1-cp311-cp311-musllinux_1_1_x86_64.whl", hash = "sha256:6a4f5ccead6d18ec072ac0b84420e95d27c1cdf5c9f1bc8fbd8daf86bd94f43d"}, + {file = "httptools-0.6.1-cp311-cp311-win_amd64.whl", hash = "sha256:5cceac09f164bcba55c0500a18fe3c47df29b62353198e4f37bbcc5d591172c3"}, + {file = "httptools-0.6.1-cp312-cp312-macosx_10_9_universal2.whl", hash = "sha256:75c8022dca7935cba14741a42744eee13ba05db00b27a4b940f0d646bd4d56d0"}, + {file = "httptools-0.6.1-cp312-cp312-macosx_10_9_x86_64.whl", hash = "sha256:48ed8129cd9a0d62cf4d1575fcf90fb37e3ff7d5654d3a5814eb3d55f36478c2"}, + {file = "httptools-0.6.1-cp312-cp312-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:6f58e335a1402fb5a650e271e8c2d03cfa7cea46ae124649346d17bd30d59c90"}, + {file = "httptools-0.6.1-cp312-cp312-manylinux_2_5_x86_64.manylinux1_x86_64.manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:93ad80d7176aa5788902f207a4e79885f0576134695dfb0fefc15b7a4648d503"}, + {file = "httptools-0.6.1-cp312-cp312-musllinux_1_1_aarch64.whl", hash = "sha256:9bb68d3a085c2174c2477eb3ffe84ae9fb4fde8792edb7bcd09a1d8467e30a84"}, + {file = "httptools-0.6.1-cp312-cp312-musllinux_1_1_x86_64.whl", hash = "sha256:b512aa728bc02354e5ac086ce76c3ce635b62f5fbc32ab7082b5e582d27867bb"}, + {file = "httptools-0.6.1-cp312-cp312-win_amd64.whl", hash = "sha256:97662ce7fb196c785344d00d638fc9ad69e18ee4bfb4000b35a52efe5adcc949"}, + {file = "httptools-0.6.1-cp38-cp38-macosx_10_9_universal2.whl", hash = "sha256:8e216a038d2d52ea13fdd9b9c9c7459fb80d78302b257828285eca1c773b99b3"}, + {file = "httptools-0.6.1-cp38-cp38-macosx_10_9_x86_64.whl", hash = "sha256:3e802e0b2378ade99cd666b5bffb8b2a7cc8f3d28988685dc300469ea8dd86cb"}, + {file = "httptools-0.6.1-cp38-cp38-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:4bd3e488b447046e386a30f07af05f9b38d3d368d1f7b4d8f7e10af85393db97"}, + {file = "httptools-0.6.1-cp38-cp38-manylinux_2_5_x86_64.manylinux1_x86_64.manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:fe467eb086d80217b7584e61313ebadc8d187a4d95bb62031b7bab4b205c3ba3"}, + {file = "httptools-0.6.1-cp38-cp38-musllinux_1_1_aarch64.whl", hash = "sha256:3c3b214ce057c54675b00108ac42bacf2ab8f85c58e3f324a4e963bbc46424f4"}, + {file = "httptools-0.6.1-cp38-cp38-musllinux_1_1_x86_64.whl", hash = "sha256:8ae5b97f690badd2ca27cbf668494ee1b6d34cf1c464271ef7bfa9ca6b83ffaf"}, + {file = "httptools-0.6.1-cp38-cp38-win_amd64.whl", hash = "sha256:405784577ba6540fa7d6ff49e37daf104e04f4b4ff2d1ac0469eaa6a20fde084"}, + {file = "httptools-0.6.1-cp39-cp39-macosx_10_9_universal2.whl", hash = "sha256:95fb92dd3649f9cb139e9c56604cc2d7c7bf0fc2e7c8d7fbd58f96e35eddd2a3"}, + {file = "httptools-0.6.1-cp39-cp39-macosx_10_9_x86_64.whl", hash = "sha256:dcbab042cc3ef272adc11220517278519adf8f53fd3056d0e68f0a6f891ba94e"}, + {file = "httptools-0.6.1-cp39-cp39-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:0cf2372e98406efb42e93bfe10f2948e467edfd792b015f1b4ecd897903d3e8d"}, + {file = "httptools-0.6.1-cp39-cp39-manylinux_2_5_x86_64.manylinux1_x86_64.manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:678fcbae74477a17d103b7cae78b74800d795d702083867ce160fc202104d0da"}, + {file = "httptools-0.6.1-cp39-cp39-musllinux_1_1_aarch64.whl", hash = "sha256:e0b281cf5a125c35f7f6722b65d8542d2e57331be573e9e88bc8b0115c4a7a81"}, + {file = "httptools-0.6.1-cp39-cp39-musllinux_1_1_x86_64.whl", hash = "sha256:95658c342529bba4e1d3d2b1a874db16c7cca435e8827422154c9da76ac4e13a"}, + {file = "httptools-0.6.1-cp39-cp39-win_amd64.whl", hash = "sha256:7ebaec1bf683e4bf5e9fbb49b8cc36da482033596a415b3e4ebab5a4c0d7ec5e"}, + {file = "httptools-0.6.1.tar.gz", hash = "sha256:c6e26c30455600b95d94b1b836085138e82f177351454ee841c148f93a9bad5a"}, +] + +[package.extras] +test = ["Cython (>=0.29.24,<0.30.0)"] + [[package]] name = "httpx" version = "0.27.0" @@ -2413,6 +2462,51 @@ typing-extensions = {version = ">=4.0", markers = "python_version < \"3.11\""} [package.extras] standard = ["colorama (>=0.4)", "httptools (>=0.5.0)", "python-dotenv (>=0.13)", "pyyaml (>=5.1)", "uvloop (>=0.14.0,!=0.15.0,!=0.15.1)", "watchfiles (>=0.13)", "websockets (>=10.4)"] +[[package]] +name = "uvloop" +version = "0.19.0" +description = "Fast implementation of asyncio event loop on top of libuv" +category = "main" +optional = false +python-versions = ">=3.8.0" +files = [ + {file = "uvloop-0.19.0-cp310-cp310-macosx_10_9_universal2.whl", hash = "sha256:de4313d7f575474c8f5a12e163f6d89c0a878bc49219641d49e6f1444369a90e"}, + {file = "uvloop-0.19.0-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:5588bd21cf1fcf06bded085f37e43ce0e00424197e7c10e77afd4bbefffef428"}, + {file = "uvloop-0.19.0-cp310-cp310-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:7b1fd71c3843327f3bbc3237bedcdb6504fd50368ab3e04d0410e52ec293f5b8"}, + {file = "uvloop-0.19.0-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:5a05128d315e2912791de6088c34136bfcdd0c7cbc1cf85fd6fd1bb321b7c849"}, + {file = "uvloop-0.19.0-cp310-cp310-musllinux_1_1_aarch64.whl", hash = "sha256:cd81bdc2b8219cb4b2556eea39d2e36bfa375a2dd021404f90a62e44efaaf957"}, + {file = "uvloop-0.19.0-cp310-cp310-musllinux_1_1_x86_64.whl", hash = "sha256:5f17766fb6da94135526273080f3455a112f82570b2ee5daa64d682387fe0dcd"}, + {file = "uvloop-0.19.0-cp311-cp311-macosx_10_9_universal2.whl", hash = "sha256:4ce6b0af8f2729a02a5d1575feacb2a94fc7b2e983868b009d51c9a9d2149bef"}, + {file = "uvloop-0.19.0-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:31e672bb38b45abc4f26e273be83b72a0d28d074d5b370fc4dcf4c4eb15417d2"}, + {file = "uvloop-0.19.0-cp311-cp311-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:570fc0ed613883d8d30ee40397b79207eedd2624891692471808a95069a007c1"}, + {file = "uvloop-0.19.0-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:5138821e40b0c3e6c9478643b4660bd44372ae1e16a322b8fc07478f92684e24"}, + {file = "uvloop-0.19.0-cp311-cp311-musllinux_1_1_aarch64.whl", hash = "sha256:91ab01c6cd00e39cde50173ba4ec68a1e578fee9279ba64f5221810a9e786533"}, + {file = "uvloop-0.19.0-cp311-cp311-musllinux_1_1_x86_64.whl", hash = "sha256:47bf3e9312f63684efe283f7342afb414eea4d3011542155c7e625cd799c3b12"}, + {file = "uvloop-0.19.0-cp312-cp312-macosx_10_9_universal2.whl", hash = "sha256:da8435a3bd498419ee8c13c34b89b5005130a476bda1d6ca8cfdde3de35cd650"}, + {file = "uvloop-0.19.0-cp312-cp312-macosx_10_9_x86_64.whl", hash = "sha256:02506dc23a5d90e04d4f65c7791e65cf44bd91b37f24cfc3ef6cf2aff05dc7ec"}, + {file = "uvloop-0.19.0-cp312-cp312-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:2693049be9d36fef81741fddb3f441673ba12a34a704e7b4361efb75cf30befc"}, + {file = "uvloop-0.19.0-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:7010271303961c6f0fe37731004335401eb9075a12680738731e9c92ddd96ad6"}, + {file = "uvloop-0.19.0-cp312-cp312-musllinux_1_1_aarch64.whl", hash = "sha256:5daa304d2161d2918fa9a17d5635099a2f78ae5b5960e742b2fcfbb7aefaa593"}, + {file = "uvloop-0.19.0-cp312-cp312-musllinux_1_1_x86_64.whl", hash = "sha256:7207272c9520203fea9b93843bb775d03e1cf88a80a936ce760f60bb5add92f3"}, + {file = "uvloop-0.19.0-cp38-cp38-macosx_10_9_universal2.whl", hash = "sha256:78ab247f0b5671cc887c31d33f9b3abfb88d2614b84e4303f1a63b46c046c8bd"}, + {file = "uvloop-0.19.0-cp38-cp38-macosx_10_9_x86_64.whl", hash = "sha256:472d61143059c84947aa8bb74eabbace30d577a03a1805b77933d6bd13ddebbd"}, + {file = "uvloop-0.19.0-cp38-cp38-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:45bf4c24c19fb8a50902ae37c5de50da81de4922af65baf760f7c0c42e1088be"}, + {file = "uvloop-0.19.0-cp38-cp38-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:271718e26b3e17906b28b67314c45d19106112067205119dddbd834c2b7ce797"}, + {file = "uvloop-0.19.0-cp38-cp38-musllinux_1_1_aarch64.whl", hash = "sha256:34175c9fd2a4bc3adc1380e1261f60306344e3407c20a4d684fd5f3be010fa3d"}, + {file = "uvloop-0.19.0-cp38-cp38-musllinux_1_1_x86_64.whl", hash = "sha256:e27f100e1ff17f6feeb1f33968bc185bf8ce41ca557deee9d9bbbffeb72030b7"}, + {file = "uvloop-0.19.0-cp39-cp39-macosx_10_9_universal2.whl", hash = "sha256:13dfdf492af0aa0a0edf66807d2b465607d11c4fa48f4a1fd41cbea5b18e8e8b"}, + {file = "uvloop-0.19.0-cp39-cp39-macosx_10_9_x86_64.whl", hash = "sha256:6e3d4e85ac060e2342ff85e90d0c04157acb210b9ce508e784a944f852a40e67"}, + {file = "uvloop-0.19.0-cp39-cp39-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:8ca4956c9ab567d87d59d49fa3704cf29e37109ad348f2d5223c9bf761a332e7"}, + {file = "uvloop-0.19.0-cp39-cp39-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:f467a5fd23b4fc43ed86342641f3936a68ded707f4627622fa3f82a120e18256"}, + {file = "uvloop-0.19.0-cp39-cp39-musllinux_1_1_aarch64.whl", hash = "sha256:492e2c32c2af3f971473bc22f086513cedfc66a130756145a931a90c3958cb17"}, + {file = "uvloop-0.19.0-cp39-cp39-musllinux_1_1_x86_64.whl", hash = "sha256:2df95fca285a9f5bfe730e51945ffe2fa71ccbfdde3b0da5772b4ee4f2e770d5"}, + {file = "uvloop-0.19.0.tar.gz", hash = "sha256:0246f4fd1bf2bf702e06b0d45ee91677ee5c31242f39aab4ea6fe0c51aedd0fd"}, +] + +[package.extras] +docs = ["Sphinx (>=4.1.2,<4.2.0)", "sphinx-rtd-theme (>=0.5.2,<0.6.0)", "sphinxcontrib-asyncio (>=0.3.0,<0.4.0)"] +test = ["Cython (>=0.29.36,<0.30.0)", "aiohttp (==3.9.0b0)", "aiohttp (>=3.8.1)", "flake8 (>=5.0,<6.0)", "mypy (>=0.800)", "psutil", "pyOpenSSL (>=23.0.0,<23.1.0)", "pycodestyle (>=2.9.0,<2.10.0)"] + [[package]] name = "websockets" version = "11.0.3" @@ -2696,4 +2790,4 @@ testing = ["big-O", "jaraco.functools", "jaraco.itertools", "more-itertools", "p [metadata] lock-version = "2.0" python-versions = "^3.8.1" -content-hash = "d2b0e968cff39082c16334498a3571c90b10fec28268f89821f13e2764990393" +content-hash = "8b93ddec633ee5bdd76f4de18d7ef4eb48b167c88983e67c88a32b0bb4170149" diff --git a/api/pyproject.toml b/api/pyproject.toml index 268072b8..87ca662c 100644 --- a/api/pyproject.toml +++ b/api/pyproject.toml @@ -28,6 +28,8 @@ psycopg = {extras = ["binary"], version = "^3.1.18"} langchain = "^0.1.12" langchain-openai = "^0.0.8" httpx = "^0.27.0" +uvloop = "^0.19.0" +httptools = "^0.6.1" [tool.ruff.lint] # from https://docs.astral.sh/ruff/linter/#rule-selection example diff --git a/api/src/agent.py b/api/src/agent.py index 8ade84b1..2cedd9d9 100644 --- a/api/src/agent.py +++ b/api/src/agent.py @@ -27,12 +27,12 @@ system_dialectic: SystemMessagePromptTemplate = SystemMessagePromptTemplate( llm: ChatOpenAI = ChatOpenAI(model_name="gpt-4") -async def chat( +async def prep_inference( + db: AsyncSession, app_id: uuid.UUID, user_id: uuid.UUID, session_id: uuid.UUID, query: str, - db: AsyncSession, ): collection = await crud.get_collection_by_name(db, app_id, user_id, "honcho") retrieved_facts = None @@ -58,6 +58,19 @@ async def chat( dialectic_prompt = ChatPromptTemplate.from_messages([system_dialectic]) chain = dialectic_prompt | llm + return (chain, retrieved_facts) + + +async def chat( + app_id: uuid.UUID, + user_id: uuid.UUID, + session_id: uuid.UUID, + query: str, + db: AsyncSession, +): + (chain, retrieved_facts) = await prep_inference( + db, app_id, user_id, session_id, query + ) response = await chain.ainvoke( { "agent_input": query, @@ -68,9 +81,19 @@ async def chat( return schemas.AgentChat(content=response.content) -async def hydrate(): - pass - - -async def insight(): - pass +async def stream( + app_id: uuid.UUID, + user_id: uuid.UUID, + session_id: uuid.UUID, + query: str, + db: AsyncSession, +): + (chain, retrieved_facts) = await prep_inference( + db, app_id, user_id, session_id, query + ) + return chain.astream( + { + "agent_input": query, + "retrieved_facts": retrieved_facts if retrieved_facts else "None", + } + ) diff --git a/api/src/deriver.py b/api/src/deriver.py index 1149708c..60d0f41f 100644 --- a/api/src/deriver.py +++ b/api/src/deriver.py @@ -34,7 +34,7 @@ if SENTRY_ENABLED: SUPABASE_ID = os.getenv("SUPABASE_ID") SUPABASE_API_KEY = os.getenv("SUPABASE_API_KEY") -llm = ChatOpenAI(model_name="gpt-3.5") +llm = ChatOpenAI(model_name="gpt-3.5-turbo") output_parser = NumberedListOutputParser() SYSTEM_DERIVE_FACTS = load_prompt( diff --git a/api/src/main.py b/api/src/main.py index 08e567f8..b018897b 100644 --- a/api/src/main.py +++ b/api/src/main.py @@ -1,18 +1,11 @@ import json import logging import os -import re -import uuid from contextlib import asynccontextmanager -from typing import Optional, Sequence -import httpx import sentry_sdk -from fastapi import ( - APIRouter, - FastAPI, - Request, -) +from fastapi import APIRouter, FastAPI +from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import PlainTextResponse from fastapi_pagination import add_pagination from opentelemetry import trace @@ -46,8 +39,6 @@ from slowapi.errors import RateLimitExceeded from slowapi.middleware import SlowAPIMiddleware from slowapi.util import get_remote_address from starlette.exceptions import HTTPException as StarletteHTTPException -from starlette.middleware.base import BaseHTTPMiddleware -from starlette.responses import Response from src.routers import ( apps, @@ -199,12 +190,43 @@ async def lifespan(app: FastAPI): await engine.dispose() -app = FastAPI(lifespan=lifespan) +app = FastAPI( + lifespan=lifespan, + servers=[ + {"url": "http://127.0.0.1:8000", "description": "Local Development Server"}, + {"url": "https:/demo.honcho.dev", "description": "Demo Server"}, + ], + title="Honcho API", + summary="An API for adding personalization to AI Apps", + description="""This API is used to store data and get insights about users for AI + applications""", + version="0.1.0", + contact={ + "name": "Plastic Labs", + "url": "https://plasticlabs.ai", + "email": "hello@plasticlabs.ai", + }, + license_info={ + "name": "GNU Affero General Public License v3.0", + "identifier": "AGPL-3.0-only", + "url": "https://github.com/plastic-labs/honcho/blob/main/LICENSE", + }, +) + +origins = ["http://localhost", "http://127.0.0.1:8000", "https://demo.honcho.dev"] + +app.add_middleware( + CORSMiddleware, + allow_origins=origins, + allow_credentials=True, + allow_methods=["*"], + allow_headers=["*"], +) + if OPENTELEMTRY_ENABLED: FastAPIInstrumentor().instrument_app(app) - router = APIRouter(prefix="/apps/{app_id}/users/{user_id}") # Create a Limiter instance @@ -221,56 +243,6 @@ app.add_middleware(SlowAPIMiddleware) add_pagination(app) -USE_AUTH_SERVICE = os.getenv("USE_AUTH_SERVICE", "False").lower() == "true" -AUTH_SERVICE_URL = os.getenv("AUTH_SERVICE_URL", "http://localhost:8001") - - -class BearerTokenMiddleware(BaseHTTPMiddleware): - async def dispatch(self, request: Request, call_next): - authorization: Optional[str] = request.headers.get("Authorization") - if authorization: - scheme, _, token = authorization.partition(" ") - if scheme.lower() == "bearer" and token: - id_pattern = r"\/apps\/([^\/]+)" - name_pattern = r"\/apps\/name\/([^\/]+)|\/apps\/get_or_create\/([^\/]+)" - match_id = re.search(id_pattern, request.url.path) - match_name = re.search(name_pattern, request.url.path) - payload = {"token": token} - if match_name: - payload["name"] = match_name.group(1) - elif match_id: - payload["app_id"] = match_id.group(1) - - res = httpx.get( - f"{AUTH_SERVICE_URL}/validate", - params=payload, - ) - data = res.json() - if ( - data["app_id"] or data["name"] - ): # Anything that checks app_id if True is valid - return await call_next(request) - if data["token"]: - check_pattern = r"^\/apps$|^\/apps\/get_or_create" - match = re.search(check_pattern, request.url.path) - if match: - return await call_next(request) - - return Response(content="Invalid token.", status_code=400) - else: - return Response( - content="Invalid authentication scheme.", status_code=400 - ) - - exclude_paths = ["/docs", "/redoc", "/openapi.json"] - if request.url.path in exclude_paths: - return await call_next(request) - return Response(content="Authorization header missing.", status_code=401) - - -if USE_AUTH_SERVICE: - app.add_middleware(BearerTokenMiddleware) - @app.exception_handler(StarletteHTTPException) async def http_exception_handler(request, exc): diff --git a/api/src/routers/apps.py b/api/src/routers/apps.py index e9d3df63..8f85972c 100644 --- a/api/src/routers/apps.py +++ b/api/src/routers/apps.py @@ -3,11 +3,14 @@ import uuid from typing import Optional import httpx -from fastapi import APIRouter, HTTPException, Request +from fastapi import APIRouter, Depends, HTTPException, Request +from psycopg.errors import UniqueViolation +from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncSession from src import crud, schemas from src.dependencies import db +from src.security import auth router = APIRouter( prefix="/apps", @@ -16,7 +19,9 @@ router = APIRouter( @router.get("/{app_id}", response_model=schemas.App) -async def get_app(request: Request, app_id: uuid.UUID, db=db): +async def get_app( + request: Request, app_id: uuid.UUID, db=db, auth: dict = Depends(auth) +): """Get an App by ID Args: @@ -33,7 +38,9 @@ async def get_app(request: Request, app_id: uuid.UUID, db=db): @router.get("/name/{name}", response_model=schemas.App) -async def get_app_by_name(request: Request, name: str, db=db): +async def get_app_by_name( + request: Request, name: str, db=db, auth: dict = Depends(auth) +): """Get an App by Name Args: @@ -50,7 +57,9 @@ async def get_app_by_name(request: Request, name: str, db=db): @router.post("", response_model=schemas.App) -async def create_app(request: Request, app: schemas.AppCreate, db=db): +async def create_app( + request: Request, app: schemas.AppCreate, db=db, auth=Depends(auth) +): """Create an App Args: @@ -60,34 +69,42 @@ async def create_app(request: Request, app: schemas.AppCreate, db=db): schemas.App: Created App object """ - USE_AUTH_SERVICE = os.getenv("USE_AUTH_SERVICE", "False").lower() == "true" - if USE_AUTH_SERVICE: - AUTH_SERVICE_URL = os.getenv("AUTH_SERVICE_URL", "http://localhost:8001") - authorization: Optional[str] = request.headers.get("Authorization") - if authorization: - scheme, _, token = authorization.partition(" ") - if token is not None: - honcho_app = await crud.create_app(db, app=app) - # if token == "default": - # return honcho_app - res = httpx.put( - f"{AUTH_SERVICE_URL}/organizations", - json={ - "id": str(honcho_app.id), - "name": honcho_app.name, - "token": token, - }, - ) - data = res.json() - if data: - return honcho_app - else: + # USE_AUTH_SERVICE = os.getenv("USE_AUTH_SERVICE", "False").lower() == "true" + # if USE_AUTH_SERVICE: + # AUTH_SERVICE_URL = os.getenv("AUTH_SERVICE_URL", "http://localhost:8001") + # authorization: Optional[str] = request.headers.get("Authorization") + # if authorization: + # scheme, _, token = authorization.partition(" ") + # if token is not None: + # honcho_app = await crud.create_app(db, app=app) + # # if token == "default": + # # return honcho_app + # res = httpx.put( + # f"{AUTH_SERVICE_URL}/organizations", + # json={ + # "id": str(honcho_app.id), + # "name": honcho_app.name, + # "token": token, + # }, + # ) + # data = res.json() + # if data: + # return honcho_app + # else: + try: honcho_app = await crud.create_app(db, app=app) return honcho_app + except IntegrityError as e: + raise HTTPException( + status_code=406, detail="App with name may already exist" + ) from e + except Exception as e: + raise HTTPException(status_code=400, detail="Unknown Error") from e + @router.get("/get_or_create/{name}", response_model=schemas.App) -async def get_or_create_app(request: Request, name: str, db=db): +async def get_or_create_app(request: Request, name: str, db=db, auth=Depends(auth)): """Get or Create an App Args: @@ -106,7 +123,11 @@ async def get_or_create_app(request: Request, name: str, db=db): @router.put("/{app_id}", response_model=schemas.App) async def update_app( - request: Request, app_id: uuid.UUID, app: schemas.AppUpdate, db=db + request: Request, + app_id: uuid.UUID, + app: schemas.AppUpdate, + db=db, + auth=Depends(auth), ): """Update an App diff --git a/api/src/routers/collections.py b/api/src/routers/collections.py index 1f9e7889..03515564 100644 --- a/api/src/routers/collections.py +++ b/api/src/routers/collections.py @@ -1,13 +1,14 @@ import json -from typing import Optional import uuid -from fastapi import APIRouter, HTTPException, Request +from typing import Optional + +from fastapi import APIRouter, Depends, HTTPException, Request from fastapi_pagination import Page from fastapi_pagination.ext.sqlalchemy import paginate from src import crud, schemas from src.dependencies import db - +from src.security import auth router = APIRouter( prefix="/apps/{app_id}/users/{user_id}/collections", @@ -23,6 +24,7 @@ async def get_collections( reverse: Optional[bool] = False, filter: Optional[str] = None, db=db, + auth=Depends(auth), ): """Get All Collections for a User @@ -64,13 +66,14 @@ async def get_collections( # return honcho_collection -@router.get("/{name}", response_model=schemas.Collection) +@router.get("/name/{name}", response_model=schemas.Collection) async def get_collection_by_name( request: Request, app_id: uuid.UUID, user_id: uuid.UUID, name: str, db=db, + auth=Depends(auth), ) -> schemas.Collection: honcho_collection = await crud.get_collection_by_name( db, app_id=app_id, user_id=user_id, name=name @@ -82,6 +85,25 @@ async def get_collection_by_name( return honcho_collection +@router.get("/{collection_id}", response_model=schemas.Collection) +async def get_collection_by_id( + request: Request, + app_id: uuid.UUID, + user_id: uuid.UUID, + collection_id: uuid.UUID, + db=db, + auth=Depends(auth), +) -> schemas.Collection: + honcho_collection = await crud.get_collection_by_id( + db, app_id=app_id, user_id=user_id, collection_id=collection_id + ) + if honcho_collection is None: + raise HTTPException( + status_code=404, detail="collection not found or does not belong to user" + ) + return honcho_collection + + @router.post("", response_model=schemas.Collection) async def create_collection( request: Request, @@ -89,6 +111,7 @@ async def create_collection( user_id: uuid.UUID, collection: schemas.CollectionCreate, db=db, + auth=Depends(auth), ): if collection.name == "honcho": raise HTTPException( @@ -114,6 +137,7 @@ async def update_collection( collection_id: uuid.UUID, collection: schemas.CollectionUpdate, db=db, + auth=Depends(auth), ): if collection.name is None: raise HTTPException( @@ -148,6 +172,7 @@ async def delete_collection( user_id: uuid.UUID, collection_id: uuid.UUID, db=db, + auth=Depends(auth), ): response = await crud.delete_collection( db, app_id=app_id, user_id=user_id, collection_id=collection_id diff --git a/api/src/routers/documents.py b/api/src/routers/documents.py index cb2b8f7b..c33864a8 100644 --- a/api/src/routers/documents.py +++ b/api/src/routers/documents.py @@ -1,13 +1,14 @@ import json -from typing import Optional, Sequence import uuid -from fastapi import APIRouter, HTTPException, Request +from typing import Optional, Sequence + +from fastapi import APIRouter, Depends, HTTPException, Request from fastapi_pagination import Page from fastapi_pagination.ext.sqlalchemy import paginate from src import crud, schemas from src.dependencies import db - +from src.security import auth router = APIRouter( prefix="/apps/{app_id}/users/{user_id}/collections/{collection_id}", @@ -24,6 +25,7 @@ async def get_documents( reverse: Optional[bool] = False, filter: Optional[str] = None, db=db, + auth=Depends(auth), ): try: data = None @@ -59,6 +61,7 @@ async def get_document( collection_id: uuid.UUID, document_id: uuid.UUID, db=db, + auth=Depends(auth), ): honcho_document = await crud.get_document( db, @@ -84,6 +87,7 @@ async def query_documents( top_k: int = 5, filter: Optional[str] = None, db=db, + auth=Depends(auth), ): if top_k is not None and top_k > 50: top_k = 50 # TODO see if we need to paginate this @@ -109,6 +113,7 @@ async def create_document( collection_id: uuid.UUID, document: schemas.DocumentCreate, db=db, + auth=Depends(auth), ): try: return await crud.create_document( @@ -136,6 +141,7 @@ async def update_document( document_id: uuid.UUID, document: schemas.DocumentUpdate, db=db, + auth=Depends(auth), ): if document.content is None and document.metadata is None: raise HTTPException( @@ -159,6 +165,7 @@ async def delete_document( collection_id: uuid.UUID, document_id: uuid.UUID, db=db, + auth=Depends(auth), ): response = await crud.delete_document( db, diff --git a/api/src/routers/messages.py b/api/src/routers/messages.py index e901dc56..72294f43 100644 --- a/api/src/routers/messages.py +++ b/api/src/routers/messages.py @@ -1,13 +1,14 @@ import json -from typing import Optional import uuid -from fastapi import APIRouter, HTTPException, Request +from typing import Optional + +from fastapi import APIRouter, Depends, HTTPException, Request from fastapi_pagination import Page from fastapi_pagination.ext.sqlalchemy import paginate from src import crud, schemas from src.dependencies import db - +from src.security import auth router = APIRouter( prefix="/apps/{app_id}/users/{user_id}/sessions/{session_id}/messages", @@ -23,6 +24,7 @@ async def create_message_for_session( session_id: uuid.UUID, message: schemas.MessageCreate, db=db, + auth=Depends(auth), ): """Adds a message to a session @@ -58,6 +60,7 @@ async def get_messages( reverse: Optional[bool] = False, filter: Optional[str] = None, db=db, + auth=Depends(auth), ): """Get all messages for a session @@ -102,6 +105,7 @@ async def get_message( session_id: uuid.UUID, message_id: uuid.UUID, db=db, + auth=Depends(auth), ): """ """ honcho_message = await crud.get_message( @@ -121,6 +125,7 @@ async def update_message( message_id: uuid.UUID, message: schemas.MessageUpdate, db=db, + auth=Depends(auth), ): """Update's the metadata of a message""" if message.metadata is None: diff --git a/api/src/routers/metamessages.py b/api/src/routers/metamessages.py index bf833290..32690ce2 100644 --- a/api/src/routers/metamessages.py +++ b/api/src/routers/metamessages.py @@ -1,13 +1,14 @@ import json -from typing import Optional import uuid -from fastapi import APIRouter, HTTPException, Request +from typing import Optional + +from fastapi import APIRouter, Depends, HTTPException, Request from fastapi_pagination import Page from fastapi_pagination.ext.sqlalchemy import paginate from src import crud, schemas from src.dependencies import db - +from src.security import auth router = APIRouter( prefix="/apps/{app_id}/users/{user_id}/sessions/{session_id}/metamessages", @@ -23,6 +24,7 @@ async def create_metamessage( session_id: uuid.UUID, metamessage: schemas.MetamessageCreate, db=db, + auth=Depends(auth), ): """Adds a message to a session @@ -64,6 +66,7 @@ async def get_metamessages( reverse: Optional[bool] = False, filter: Optional[str] = None, db=db, + auth=Depends(auth), ): """Get all messages for a session @@ -114,6 +117,7 @@ async def get_metamessage( message_id: uuid.UUID, metamessage_id: uuid.UUID, db=db, + auth=Depends(auth), ): """Get a specific Metamessage by ID @@ -154,6 +158,7 @@ async def update_metamessage( metamessage_id: uuid.UUID, metamessage: schemas.MetamessageUpdate, db=db, + auth=Depends(auth), ): """Update's the metadata of a metamessage""" if metamessage.metadata is None: diff --git a/api/src/routers/sessions.py b/api/src/routers/sessions.py index af17e725..5b831398 100644 --- a/api/src/routers/sessions.py +++ b/api/src/routers/sessions.py @@ -2,12 +2,14 @@ import json import uuid from typing import Optional -from fastapi import APIRouter, HTTPException, Request +from fastapi import APIRouter, Depends, HTTPException, Request +from fastapi.responses import StreamingResponse from fastapi_pagination import Page from fastapi_pagination.ext.sqlalchemy import paginate from src import agent, crud, schemas from src.dependencies import db +from src.security import auth router = APIRouter( prefix="/apps/{app_id}/users/{user_id}/sessions", @@ -25,6 +27,7 @@ async def get_sessions( reverse: Optional[bool] = False, filter: Optional[str] = None, db=db, + auth=Depends(auth), ): """Get All Sessions for a User @@ -65,6 +68,7 @@ async def create_session( user_id: uuid.UUID, session: schemas.SessionCreate, db=db, + auth=Depends(auth), ): """Create a Session for a User @@ -93,6 +97,7 @@ async def update_session( session_id: uuid.UUID, session: schemas.SessionUpdate, db=db, + auth=Depends(auth), ): """Update the metadata of a Session @@ -124,6 +129,7 @@ async def delete_session( user_id: uuid.UUID, session_id: uuid.UUID, db=db, + auth=Depends(auth), ): """Delete a session by marking it as inactive @@ -156,6 +162,7 @@ async def get_session( user_id: uuid.UUID, session_id: uuid.UUID, db=db, + auth=Depends(auth), ): """Get a specific session for a user by ID @@ -187,7 +194,41 @@ async def get_chat( session_id: uuid.UUID, query: str, db=db, + auth=Depends(auth), ): return await agent.chat( app_id=app_id, user_id=user_id, session_id=session_id, query=query, db=db ) + + +@router.get( + "/{session_id}/chat/stream", + responses={ + 200: { + "description": "Chat stream", + "content": { + "text/event-stream": {"schema": {"type": "string", "format": "binary"}} + }, + } + }, +) +async def get_chat_stream( + request: Request, + app_id: uuid.UUID, + user_id: uuid.UUID, + session_id: uuid.UUID, + query: str, + db=db, + auth=Depends(auth), +): + return StreamingResponse( + await agent.stream( + app_id=app_id, + user_id=user_id, + session_id=session_id, + query=query, + db=db, + ), + media_type="text/event-stream", + status_code=200, + ) diff --git a/api/src/routers/users.py b/api/src/routers/users.py index c79c2340..095d6f27 100644 --- a/api/src/routers/users.py +++ b/api/src/routers/users.py @@ -1,12 +1,15 @@ import json -from typing import Optional import uuid -from fastapi import APIRouter, Request +from typing import Optional + +from fastapi import APIRouter, Depends, HTTPException, Request from fastapi_pagination import Page from fastapi_pagination.ext.sqlalchemy import paginate +from sqlalchemy.exc import IntegrityError from src import crud, schemas from src.dependencies import db +from src.security import auth router = APIRouter( prefix="/apps/{app_id}/users", @@ -20,6 +23,7 @@ async def create_user( app_id: uuid.UUID, user: schemas.UserCreate, db=db, + auth=Depends(auth), ): """Create a User @@ -33,7 +37,12 @@ async def create_user( """ print("running create_user") - return await crud.create_user(db, app_id=app_id, user=user) + try: + return await crud.create_user(db, app_id=app_id, user=user) + except IntegrityError as e: + raise HTTPException( + status_code=406, detail="User with name may already exist" + ) from e @router.get("", response_model=Page[schemas.User]) @@ -43,6 +52,7 @@ async def get_users( reverse: bool = False, filter: Optional[str] = None, db=db, + auth=Depends(auth), ): """Get All Users for an App @@ -63,12 +73,13 @@ async def get_users( ) -@router.get("/{name}", response_model=schemas.User) +@router.get("/name/{name}", response_model=schemas.User) async def get_user_by_name( request: Request, app_id: uuid.UUID, name: str, db=db, + auth=Depends(auth), ): """Get a User @@ -84,8 +95,32 @@ async def get_user_by_name( return await crud.get_user_by_name(db, app_id=app_id, name=name) +@router.get("/{user_id}", response_model=schemas.User) +async def get_user( + request: Request, + app_id: uuid.UUID, + user_id: uuid.UUID, + db=db, + auth=Depends(auth), +): + """Get a User + + Args: + app_id (uuid.UUID): The ID of the app representing the client application using + honcho + user_id (str): The User ID representing the user, managed by the user + + Returns: + schemas.User: User object + + """ + return await crud.get_user(db, app_id=app_id, user_id=user_id) + + @router.get("/get_or_create/{name}", response_model=schemas.User) -async def get_or_create_user(request: Request, app_id: uuid.UUID, name: str, db=db): +async def get_or_create_user( + request: Request, app_id: uuid.UUID, name: str, db=db, auth=Depends(auth) +): """Get or Create a User Args: @@ -99,8 +134,8 @@ async def get_or_create_user(request: Request, app_id: uuid.UUID, name: str, db= """ user = await crud.get_user_by_name(db, app_id=app_id, name=name) if user is None: - user = await crud.create_user( - db, app_id=app_id, user=schemas.UserCreate(name=name) + user = await create_user( + request=request, db=db, app_id=app_id, user=schemas.UserCreate(name=name) ) return user @@ -112,6 +147,7 @@ async def update_user( user_id: uuid.UUID, user: schemas.UserUpdate, db=db, + auth=Depends(auth), ): """Update a User diff --git a/api/src/schemas.py b/api/src/schemas.py index 1f222d93..1f9b2d63 100644 --- a/api/src/schemas.py +++ b/api/src/schemas.py @@ -52,6 +52,7 @@ class UserUpdate(UserBase): class User(UserBase): id: uuid.UUID + name: str app_id: uuid.UUID created_at: datetime.datetime h_metadata: dict = Field(exclude=True) diff --git a/api/src/security.py b/api/src/security.py new file mode 100644 index 00000000..f54758a6 --- /dev/null +++ b/api/src/security.py @@ -0,0 +1,33 @@ +import os +from typing import Annotated + +import httpx +from fastapi import APIRouter, Depends, FastAPI, HTTPException, Request +from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer + +USE_AUTH_SERVICE = os.getenv("USE_AUTH_SERVICE", "False").lower() == "true" +AUTH_SERVICE_URL = os.getenv("AUTH_SERVICE_URL", "http://localhost:8001") + +security = HTTPBearer( + auto_error=False, +) + + +async def auth( + credentials: Annotated[HTTPAuthorizationCredentials, Depends(security)], +): + if not USE_AUTH_SERVICE: + print("Test of Auth") + return True + print(credentials) + if not credentials or credentials.credentials != "test": + raise HTTPException(status_code=401, detail="Invalid access token") + # payload = {"token": token} + # res = httpx.get( + # f"{AUTH_SERVICE_URL}/validate", + # params=payload, + # ) + # data = res.json() + return {"message": "OK"} + + # return {"scheme": credentials.scheme, "token": credentials.credentials} diff --git a/sdk/honcho/client.py b/sdk/honcho/client.py index efae6ec4..082d51a4 100644 --- a/sdk/honcho/client.py +++ b/sdk/honcho/client.py @@ -449,7 +449,7 @@ class AsyncHoncho: Returns: AsyncUser: The User object """ - url = f"{self.base_url}/users/{name}" + url = f"{self.base_url}/users/name/{name}" response = await self.client.get(url) response.raise_for_status() data = response.json() @@ -793,7 +793,7 @@ class AsyncUser: AsyncCollection: The Session object of the requested Session """ - url = f"{self.base_url}/collections/{name}" + url = f"{self.base_url}/collections/name/{name}" response = await self.honcho.client.get(url) response.raise_for_status() data = response.json() diff --git a/sdk/honcho/sync_client.py b/sdk/honcho/sync_client.py index 07558110..0c3d54f8 100644 --- a/sdk/honcho/sync_client.py +++ b/sdk/honcho/sync_client.py @@ -449,7 +449,7 @@ class Honcho: Returns: User: The User object """ - url = f"{self.base_url}/users/{name}" + url = f"{self.base_url}/users/name/{name}" response = self.client.get(url) response.raise_for_status() data = response.json() @@ -793,7 +793,7 @@ class User: Collection: The Session object of the requested Session """ - url = f"{self.base_url}/collections/{name}" + url = f"{self.base_url}/collections/name/{name}" response = self.honcho.client.get(url) response.raise_for_status() data = response.json() diff --git a/sdk/tests/test_sync.py b/sdk/tests/test_sync.py index fc20b116..0077ff00 100644 --- a/sdk/tests/test_sync.py +++ b/sdk/tests/test_sync.py @@ -251,7 +251,9 @@ def test_paginated_messages(): created_session.create_message(is_user=False, content="Hi") page_size = 7 - get_message_response = created_session.get_messages(page=1, page_size=page_size) + get_message_response = created_session.get_messages( + page=1, page_size=page_size + ) assert get_message_response is not None assert isinstance(get_message_response, GetMessagePage) @@ -446,7 +448,9 @@ def test_collection_query(): collection = user.create_collection(col_name) # Add documents - doc1 = collection.create_document(content="The user loves puppies", metadata={}) + doc1 = collection.create_document( + content="The user loves puppies", metadata={} + ) doc2 = collection.create_document(content="The user owns a dog", metadata={}) doc3 = collection.create_document(content="The user is a doctor", metadata={})