[0.0.10] Release

* chore: Save Point

* fix(db): change connection logic and remove unnecessary refreshes

* fix(documents): switch to Azure embedding model

* feat: Setup pytest fixtures

* feat(tests): Initial test routes for tranche 1

* feat(test) tranche 2 of tests and associated bug fixes

* feat(test) tranche 3 of tests and associated bug fixes

* fix(tests) Address PR comments and update version and changelog
This commit is contained in:
Vineeth Voruganti 2024-07-25 13:45:27 -04:00 committed by GitHub
parent 659b13a10f
commit ddcde6bcea
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
27 changed files with 1611 additions and 688 deletions

View File

@ -4,6 +4,26 @@ All notable changes to this project will be documented in this file.
The format is based on [Keep a Changelog](http://keepachangelog.com/) The format is based on [Keep a Changelog](http://keepachangelog.com/)
and this project adheres to [Semantic Versioning](http://semver.org/). and this project adheres to [Semantic Versioning](http://semver.org/).
## [0.0.10] — 2024-07-23
### Added
* Test cases for Storage API
* Sentry tracing and profiling
* Additional Error handling
### Changed
* Document API uses same embedding endpoint as deriver
* CRUD operations use one less database call by removing extra refresh
* Use database for timestampz rather than API
* Pydantic schemas to use modern syntax
### Fixed
* Deriver queue resolution
## [0.0.9] — 2024-05-16 ## [0.0.9] — 2024-05-16
### Added ### Added

View File

@ -1,5 +1,5 @@
# 🫡 Honcho # 🫡 Honcho
![Static Badge](https://img.shields.io/badge/Version-0.0.9-blue) ![Static Badge](https://img.shields.io/badge/Version-0.0.10-blue)
[![Discord](https://img.shields.io/discord/1016845111637839922?style=flat&logo=discord&logoColor=23ffffff&label=Plastic%20Labs&labelColor=235865F2)](https://discord.gg/plasticlabs) [![Discord](https://img.shields.io/discord/1016845111637839922?style=flat&logo=discord&logoColor=23ffffff&label=Plastic%20Labs&labelColor=235865F2)](https://discord.gg/plasticlabs)
[![arXiv](https://img.shields.io/badge/arXiv-2310.06983-b31b1b.svg)](https://arxiv.org/abs/2310.06983) [![arXiv](https://img.shields.io/badge/arXiv-2310.06983-b31b1b.svg)](https://arxiv.org/abs/2310.06983)
![GitHub License](https://img.shields.io/github/license/plastic-labs/honcho) ![GitHub License](https://img.shields.io/github/license/plastic-labs/honcho)

View File

@ -1,6 +1,7 @@
version: "3.8" version: "3.8"
services: services:
api: api:
image: honcho:latest
build: build:
context: . context: .
dockerfile: Dockerfile dockerfile: Dockerfile

1114
poetry.lock generated

File diff suppressed because it is too large Load Diff

View File

@ -1,6 +1,6 @@
[tool.poetry] [tool.poetry]
name = "honcho" name = "honcho"
version = "0.0.9" version = "0.0.10"
description = "Honcho Server" description = "Honcho Server"
authors = ["Plastic Labs <hello@plasticlabs.ai>"] authors = ["Plastic Labs <hello@plasticlabs.ai>"]
readme = "README.md" readme = "README.md"
@ -8,19 +8,16 @@ readme = "README.md"
[tool.poetry.dependencies] [tool.poetry.dependencies]
python = "^3.9" python = "^3.9"
fastapi = "^0.111.0" fastapi = "^0.111.0"
uvicorn = "^0.29.0"
python-dotenv = "^1.0.0" python-dotenv = "^1.0.0"
sqlalchemy = "^2.0.30" sqlalchemy = "^2.0.30"
slowapi = "^0.1.9" slowapi = "^0.1.9"
fastapi-pagination = "^0.12.24" fastapi-pagination = "^0.12.24"
pgvector = "^0.2.5" pgvector = "^0.2.5"
openai = "^1.12.0" openai = "^1.12.0"
sentry-sdk = "^2.3.0" sentry-sdk = {extras = ["fastapi", "sqlalchemy"], version = "^2.3.1"}
greenlet = "^3.0.3" greenlet = "^3.0.3"
psycopg = {extras= ["binary"], version="^3.1.19"} psycopg = {extras= ["binary"], version="^3.1.19"}
httpx = "^0.27.0" httpx = "^0.27.0"
uvloop = "^0.19.0"
httptools = "^0.6.1"
mirascope = "^0.15.1" mirascope = "^0.15.1"
opentelemetry-instrumentation-fastapi = "^0.45b0" opentelemetry-instrumentation-fastapi = "^0.45b0"
opentelemetry-sdk = "^1.24.0" opentelemetry-sdk = "^1.24.0"
@ -28,6 +25,13 @@ opentelemetry-exporter-otlp = "^1.24.0"
opentelemetry-instrumentation-sqlalchemy = "^0.45b0" opentelemetry-instrumentation-sqlalchemy = "^0.45b0"
opentelemetry-instrumentation-logging = "^0.45b0" opentelemetry-instrumentation-logging = "^0.45b0"
[tool.poetry.group.test.dependencies]
pytest = "^8.2.2"
sqlalchemy-utils = "^0.41.2"
pytest-asyncio = "^0.23.7"
coverage = "^7.6.0"
interrogate = "^1.7.0"
[tool.ruff.lint] [tool.ruff.lint]
# from https://docs.astral.sh/ruff/linter/#rule-selection example # from https://docs.astral.sh/ruff/linter/#rule-selection example
select = [ select = [
@ -51,3 +55,6 @@ extend-immutable-calls = ["fastapi.Depends"]
[build-system] [build-system]
requires = ["poetry-core"] requires = ["poetry-core"]
build-backend = "poetry.core.masonry.api" build-backend = "poetry.core.masonry.api"
[tool.lpytest.ini_options]
asyncio_mode = "auto"

View File

@ -1,8 +1,9 @@
import datetime import datetime
import os
import uuid import uuid
from typing import Optional, Sequence from typing import Optional, Sequence
from openai import OpenAI from openai import AzureOpenAI, OpenAI
from sqlalchemy import Select, select from sqlalchemy import Select, select
from sqlalchemy.exc import IntegrityError from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
@ -10,7 +11,11 @@ from sqlalchemy.ext.asyncio import AsyncSession
# from sqlalchemy.orm import Session # from sqlalchemy.orm import Session
from . import models, schemas from . import models, schemas
openai_client = OpenAI() openai_client = AzureOpenAI(
api_key=os.getenv("AZURE_OPENAI_API_KEY"),
api_version=os.getenv("AZURE_OPENAI_API_VERSION"),
azure_endpoint=os.getenv("AZURE_OPENAI_ENDPOINT"),
)
######################################################## ########################################################
# app methods # app methods
@ -39,7 +44,7 @@ async def create_app(db: AsyncSession, app: schemas.AppCreate) -> models.App:
honcho_app = models.App(name=app.name, h_metadata=app.metadata) honcho_app = models.App(name=app.name, h_metadata=app.metadata)
db.add(honcho_app) db.add(honcho_app)
await db.commit() await db.commit()
await db.refresh(honcho_app) # await db.refresh(honcho_app)
return honcho_app return honcho_app
@ -50,12 +55,12 @@ async def update_app(
if honcho_app is None: if honcho_app is None:
raise ValueError("App not found") raise ValueError("App not found")
if app.name is not None: if app.name is not None:
honcho_app.content = app.name honcho_app.name = app.name
if app.metadata is not None: if app.metadata is not None:
honcho_app.h_metadata = app.metadata honcho_app.h_metadata = app.metadata
await db.commit() await db.commit()
await db.refresh(honcho_app) # await db.refresh(honcho_app)
return honcho_app return honcho_app
@ -83,7 +88,7 @@ async def create_user(
) )
db.add(honcho_user) db.add(honcho_user)
await db.commit() await db.commit()
await db.refresh(honcho_user) # await db.refresh(honcho_user)
return honcho_user return honcho_user
@ -139,12 +144,12 @@ async def update_user(
if honcho_user is None: if honcho_user is None:
raise ValueError("User not found") raise ValueError("User not found")
if user.name is not None: if user.name is not None:
honcho_user.content = user.name honcho_user.name = user.name
if user.metadata is not None: if user.metadata is not None:
honcho_user.h_metadata = user.metadata honcho_user.h_metadata = user.metadata
await db.commit() await db.commit()
await db.refresh(honcho_user) # await db.refresh(honcho_user)
return honcho_user return honcho_user
@ -228,8 +233,18 @@ async def create_session(
h_metadata=session.metadata, h_metadata=session.metadata,
) )
db.add(honcho_session) db.add(honcho_session)
# print("====== Testing State of ORM Object ====")
# print(honcho_session)
# print("=======================================")
#
# await db.flush()
#
# print("====== Testing State of ORM Object ====")
# print(honcho_session)
# print("=======================================")
await db.commit() await db.commit()
await db.refresh(honcho_session) # await db.refresh(honcho_session)
return honcho_session return honcho_session
@ -250,7 +265,7 @@ async def update_session(
): # Need to explicitly be there won't make it empty by default ): # Need to explicitly be there won't make it empty by default
honcho_session.h_metadata = session.metadata honcho_session.h_metadata = session.metadata
await db.commit() await db.commit()
await db.refresh(honcho_session) # await db.refresh(honcho_session)
return honcho_session return honcho_session
@ -300,7 +315,7 @@ async def create_message(
db.add(honcho_message) db.add(honcho_message)
await db.commit() await db.commit()
# await db.refresh(honcho_message, attribute_names=["id", "content", "h_metadata"]) # await db.refresh(honcho_message, attribute_names=["id", "content", "h_metadata"])
await db.refresh(honcho_message) # await db.refresh(honcho_message)
return honcho_message return honcho_message
@ -372,7 +387,7 @@ async def update_message(
): # Need to explicitly be there won't make it empty by default ): # Need to explicitly be there won't make it empty by default
honcho_message.h_metadata = message.metadata honcho_message.h_metadata = message.metadata
await db.commit() await db.commit()
await db.refresh(honcho_message) # await db.refresh(honcho_message)
return honcho_message return honcho_message
@ -388,7 +403,7 @@ async def create_metamessage(
user_id: uuid.UUID, user_id: uuid.UUID,
session_id: uuid.UUID, session_id: uuid.UUID,
): ):
message = get_message( message = await get_message(
db, db,
app_id=app_id, app_id=app_id,
session_id=session_id, session_id=session_id,
@ -407,7 +422,7 @@ async def create_metamessage(
db.add(honcho_metamessage) db.add(honcho_metamessage)
await db.commit() await db.commit()
await db.refresh(honcho_metamessage) # await db.refresh(honcho_metamessage)
return honcho_metamessage return honcho_metamessage
@ -498,7 +513,7 @@ async def update_metamessage(
if metamessage.metamessage_type is not None: if metamessage.metamessage_type is not None:
honcho_metamessage.metamessage_type = metamessage.metamessage_type honcho_metamessage.metamessage_type = metamessage.metamessage_type
await db.commit() await db.commit()
await db.refresh(honcho_metamessage) # await db.refresh(honcho_metamessage)
return honcho_metamessage return honcho_metamessage
@ -582,7 +597,7 @@ async def create_collection(
except IntegrityError: except IntegrityError:
await db.rollback() await db.rollback()
raise ValueError("Collection already exists") from None raise ValueError("Collection already exists") from None
await db.refresh(honcho_collection) # await db.refresh(honcho_collection)
return honcho_collection return honcho_collection
@ -601,12 +616,13 @@ async def update_collection(
if collection.metadata is not None: if collection.metadata is not None:
honcho_collection.h_metadata = collection.metadata honcho_collection.h_metadata = collection.metadata
try: try:
honcho_collection.name = collection.name if collection.name is not None:
await db.commit() honcho_collection.name = collection.name
await db.commit()
except IntegrityError: except IntegrityError:
await db.rollback() await db.rollback()
raise ValueError("Collection already exists") from None raise ValueError("Collection already exists") from None
await db.refresh(honcho_collection) # await db.refresh(honcho_collection)
return honcho_collection return honcho_collection
@ -700,7 +716,7 @@ async def query_documents(
top_k: int = 5, top_k: int = 5,
) -> Sequence[models.Document]: ) -> Sequence[models.Document]:
response = openai_client.embeddings.create( response = openai_client.embeddings.create(
input=query, model="text-embedding-3-small" input=query, model=os.getenv("AZURE_OPENAI_EMBED_DEPLOYMENT")
) )
embedding_query = response.data[0].embedding embedding_query = response.data[0].embedding
stmt = ( stmt = (
@ -736,7 +752,7 @@ async def create_document(
raise ValueError("Session not found or does not belong to user") raise ValueError("Session not found or does not belong to user")
response = openai_client.embeddings.create( response = openai_client.embeddings.create(
input=document.content, model="text-embedding-3-small" input=document.content, model=os.getenv("AZURE_OPENAI_EMBED_DEPLOYMENT")
) )
embedding = response.data[0].embedding embedding = response.data[0].embedding
@ -749,7 +765,7 @@ async def create_document(
) )
db.add(honcho_document) db.add(honcho_document)
await db.commit() await db.commit()
await db.refresh(honcho_document) # await db.refresh(honcho_document)
return honcho_document return honcho_document
@ -773,7 +789,7 @@ async def update_document(
if document.content is not None: if document.content is not None:
honcho_document.content = document.content honcho_document.content = document.content
response = openai_client.embeddings.create( response = openai_client.embeddings.create(
input=document.content, model="text-embedding-3-small" input=document.content, model=os.getenv("AZURE_OPENAI_EMBED_DEPLOYMENT")
) )
embedding = response.data[0].embedding embedding = response.data[0].embedding
honcho_document.embedding = embedding honcho_document.embedding = embedding
@ -782,7 +798,7 @@ async def update_document(
if document.metadata is not None: if document.metadata is not None:
honcho_document.h_metadata = document.metadata honcho_document.h_metadata = document.metadata
await db.commit() await db.commit()
await db.refresh(honcho_document) # await db.refresh(honcho_document)
return honcho_document return honcho_document

View File

@ -1,9 +1,11 @@
import os import os
from dotenv import load_dotenv from dotenv import load_dotenv
from sqlalchemy import create_engine from sqlalchemy import MetaData, create_engine
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
from sqlalchemy.ext.declarative import declarative_base from sqlalchemy.orm import declarative_base
# from sqlalchemy.ext.declarative import declarative_base
load_dotenv() load_dotenv()
@ -21,8 +23,15 @@ engine = create_async_engine(
pool_pre_ping=True, pool_pre_ping=True,
) )
SessionLocal = async_sessionmaker(autocommit=False, autoflush=False, bind=engine) SessionLocal = async_sessionmaker(
Base = declarative_base() autocommit=False, autoflush=False, expire_on_commit=False, bind=engine
)
table_schema = os.getenv("DATABASE_SCHEMA")
meta = MetaData()
if table_schema:
meta.schema = table_schema
Base = declarative_base(metadata=meta)
def scaffold_db(): def scaffold_db():

View File

@ -8,6 +8,7 @@ from typing import List
import sentry_sdk import sentry_sdk
import uvloop import uvloop
from dotenv import load_dotenv from dotenv import load_dotenv
from sentry_sdk.integrations.asyncio import AsyncioIntegration
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
@ -23,14 +24,6 @@ from .voe import (
load_dotenv() load_dotenv()
SENTRY_ENABLED = os.getenv("SENTRY_ENABLED", "False").lower() == "true"
if SENTRY_ENABLED:
sentry_sdk.init(
dsn=os.getenv("SENTRY_DSN"),
enable_tracing=True,
)
# Turn of SQLAlchemy Echo logging # Turn of SQLAlchemy Echo logging
logging.getLogger("sqlalchemy.engine.Engine").disabled = True logging.getLogger("sqlalchemy.engine.Engine").disabled = True
@ -148,7 +141,7 @@ async def process_ai_message(
) )
if user_prediction_thought_revision_response.content == "None": if user_prediction_thought_revision_response.content == "None":
print(f"\033[94mModel predicted no changes to the user prediction thought") print("\033[94mModel predicted no changes to the user prediction thought")
await add_metamessage( await add_metamessage(
db, db,
message_id, message_id,
@ -178,27 +171,27 @@ async def process_ai_message(
await db.commit() await db.commit()
# debugging # debugging
print(f"\033[94m=================") print("\033[94m=================")
print(f"\033[94mUser Prediction Thought Prompt:") print("\033[94mUser Prediction Thought Prompt:")
content_lines = str(user_prediction_thought).split("\n") content_lines = str(user_prediction_thought).split("\n")
for line in content_lines: for line in content_lines:
print(f"\033[94m{line}") print(f"\033[94m{line}")
print(f"\033[94mUser Prediction Thought:") print("\033[94mUser Prediction Thought:")
content_lines = str(user_prediction_thought_response.content).split("\n") content_lines = str(user_prediction_thought_response.content).split("\n")
for line in content_lines: for line in content_lines:
print(f"\033[94m{line}") print(f"\033[94m{line}")
print(f"\033[94m=================\033[0m") print("\033[94m=================\033[0m")
print(f"\033[95m=================") print("\033[95m=================")
print(f"\033[95mUser Prediction Thought Revision:") print("\033[95mUser Prediction Thought Revision:")
content_lines = str(user_prediction_thought_revision).split("\n") content_lines = str(user_prediction_thought_revision).split("\n")
for line in content_lines: for line in content_lines:
print(f"\033[95m{line}") print(f"\033[95m{line}")
print(f"\033[95mUser Prediction Thought Revision Response:") print("\033[95mUser Prediction Thought Revision Response:")
content_lines = str(user_prediction_thought_revision_response.content).split("\n") content_lines = str(user_prediction_thought_revision_response.content).split("\n")
for line in content_lines: for line in content_lines:
print(f"\033[95m{line}") print(f"\033[95m{line}")
print(f"\033[95m=================\033[0m") print("\033[95m=================\033[0m")
async def process_user_message( async def process_user_message(
@ -222,9 +215,10 @@ async def process_user_message(
messages_stmt = ( messages_stmt = (
select(models.Message) select(models.Message)
.where(models.Message.created_at < subquery)
.where(models.Message.session_id == session_id) .where(models.Message.session_id == session_id)
.where(models.Message.is_user == False) .where(models.Message.is_user == False)
.order_by(models.Message.created_at.desc())
.where(models.Message.created_at < subquery)
.limit(1) .limit(1)
) )
@ -265,31 +259,31 @@ async def process_user_message(
voe_derive_facts_response = await voe_derive_facts.call_async() voe_derive_facts_response = await voe_derive_facts.call_async()
# debugging # debugging
print(f"\033[93m=================") print("\033[93m=================")
print(f"\033[93mVoe Thought Prompt:") print("\033[93mVoe Thought Prompt:")
content_lines = str(voe_thought).split("\n") content_lines = str(voe_thought).split("\n")
for line in content_lines: for line in content_lines:
print(f"\033[93m{line}") print(f"\033[93m{line}")
print(f"\033[93mVoe Thought:") print("\033[93mVoe Thought:")
content_lines = str(voe_thought_response.content).split("\n") content_lines = str(voe_thought_response.content).split("\n")
for line in content_lines: for line in content_lines:
print(f"\033[93m{line}") print(f"\033[93m{line}")
print(f"\033[93m=================\033[0m") print("\033[93m=================\033[0m")
print(f"\033[93m=================") print("\033[93m=================")
print(f"\033[93mVoe Derive Facts Prompt:") print("\033[93mVoe Derive Facts Prompt:")
content_lines = str(voe_derive_facts).split("\n") content_lines = str(voe_derive_facts).split("\n")
for line in content_lines: for line in content_lines:
print(f"\033[93m{line}") print(f"\033[93m{line}")
print(f"\033[93mVoe Derive Facts Response:") print("\033[93mVoe Derive Facts Response:")
content_lines = str(voe_derive_facts_response.content).split("\n") content_lines = str(voe_derive_facts_response.content).split("\n")
for line in content_lines: for line in content_lines:
print(f"\033[93m{line}") print(f"\033[93m{line}")
print(f"\033[93m=================\033[0m") print("\033[93m=================\033[0m")
facts = re.findall(r"\d+\.\s([^\n]+)", voe_derive_facts_response.content) facts = re.findall(r"\d+\.\s([^\n]+)", voe_derive_facts_response.content)
print(f"\033[93m=================") print("\033[93m=================")
print(f"\033[93mThe Facts Themselves:") print("\033[93mThe Facts Themselves:")
print(facts) print(facts)
new_facts = await check_dups(app_id, user_id, collection_id, facts) new_facts = await check_dups(app_id, user_id, collection_id, facts)
@ -305,9 +299,9 @@ async def process_user_message(
) )
print(f"\033[93mReturned Document: {doc.content}") print(f"\033[93mReturned Document: {doc.content}")
else: else:
raise Exception(f"\033[91mUser Thought Prediction Revision NOT READY YET") raise Exception("\033[91mUser Thought Prediction Revision NOT READY YET")
else: else:
print(f"\033[91mNo AI message before this user message") print("\033[91mNo AI message before this user message")
return return
@ -362,7 +356,7 @@ async def dequeue(semaphore: asyncio.Semaphore, queue_empty_flag: asyncio.Event)
try: try:
result = await db.execute( result = await db.execute(
select(models.QueueItem) select(models.QueueItem)
.order_by(models.QueueItem.created_at) .order_by(models.QueueItem.id)
.where(models.QueueItem.processed == False) .where(models.QueueItem.processed == False)
.with_for_update(skip_locked=True) .with_for_update(skip_locked=True)
.limit(1) .limit(1)
@ -402,6 +396,17 @@ async def polling_loop(semaphore: asyncio.Semaphore, queue_empty_flag: asyncio.E
async def main(): async def main():
SENTRY_ENABLED = os.getenv("SENTRY_ENABLED", "False").lower() == "true"
if SENTRY_ENABLED:
sentry_sdk.init(
dsn=os.getenv("SENTRY_DSN"),
enable_tracing=True,
traces_sample_rate=1.0,
profiles_sample_rate=1.0,
integrations=[
AsyncioIntegration(),
],
)
semaphore = asyncio.Semaphore(1) # Limit to 5 concurrent dequeuing operations semaphore = asyncio.Semaphore(1) # Limit to 5 concurrent dequeuing operations
queue_empty_flag = asyncio.Event() # Event to signal when the queue is empty queue_empty_flag = asyncio.Event() # Event to signal when the queue is empty
await polling_loop(semaphore, queue_empty_flag) await polling_loop(semaphore, queue_empty_flag)

View File

@ -180,6 +180,8 @@ if SENTRY_ENABLED:
sentry_sdk.init( sentry_sdk.init(
dsn=os.getenv("SENTRY_DSN"), dsn=os.getenv("SENTRY_DSN"),
enable_tracing=True, enable_tracing=True,
traces_sample_rate=0.4,
profiles_sample_rate=0.4,
) )
@ -223,17 +225,6 @@ app.add_middleware(
allow_headers=["*"], allow_headers=["*"],
) )
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: if OPENTELEMTRY_ENABLED:
FastAPIInstrumentor().instrument_app(app) FastAPIInstrumentor().instrument_app(app)

View File

@ -1,5 +1,4 @@
import datetime import datetime
import os
import uuid import uuid
from dotenv import load_dotenv from dotenv import load_dotenv
@ -17,6 +16,7 @@ from sqlalchemy import (
) )
from sqlalchemy.dialects.postgresql import JSONB from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.orm import Mapped, mapped_column, relationship from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.sql import func
from .db import Base from .db import Base
@ -35,7 +35,7 @@ class App(Base):
name: Mapped[str] = mapped_column(String(512), index=True, unique=True) name: Mapped[str] = mapped_column(String(512), index=True, unique=True)
users = relationship("User", back_populates="app") users = relationship("User", back_populates="app")
created_at: Mapped[datetime.datetime] = mapped_column( created_at: Mapped[datetime.datetime] = mapped_column(
DateTime(timezone=True), default=datetime.datetime.utcnow DateTime(timezone=True), default=func.now()
) )
h_metadata: Mapped[dict] = mapped_column("metadata", JSONB, default={}) h_metadata: Mapped[dict] = mapped_column("metadata", JSONB, default={})
# Add any additional fields for an app here # Add any additional fields for an app here
@ -49,7 +49,7 @@ class User(Base):
name: Mapped[str] = mapped_column(String(512), index=True) name: Mapped[str] = mapped_column(String(512), index=True)
h_metadata: Mapped[dict] = mapped_column("metadata", JSONB, default={}) h_metadata: Mapped[dict] = mapped_column("metadata", JSONB, default={})
created_at: Mapped[datetime.datetime] = mapped_column( created_at: Mapped[datetime.datetime] = mapped_column(
DateTime(timezone=True), default=datetime.datetime.utcnow DateTime(timezone=True), default=func.now()
) )
app_id: Mapped[uuid.UUID] = mapped_column(ForeignKey("apps.id"), index=True) app_id: Mapped[uuid.UUID] = mapped_column(ForeignKey("apps.id"), index=True)
app = relationship("App", back_populates="users") app = relationship("App", back_populates="users")
@ -71,14 +71,14 @@ class Session(Base):
is_active: Mapped[bool] = mapped_column(default=True) is_active: Mapped[bool] = mapped_column(default=True)
h_metadata: Mapped[dict] = mapped_column("metadata", JSONB, default={}) h_metadata: Mapped[dict] = mapped_column("metadata", JSONB, default={})
created_at: Mapped[datetime.datetime] = mapped_column( created_at: Mapped[datetime.datetime] = mapped_column(
DateTime(timezone=True), default=datetime.datetime.utcnow DateTime(timezone=True), default=func.now()
) )
messages = relationship("Message", back_populates="session") messages = relationship("Message", back_populates="session")
user_id: Mapped[uuid.UUID] = mapped_column(ForeignKey("users.id"), index=True) user_id: Mapped[uuid.UUID] = mapped_column(ForeignKey("users.id"), index=True)
user = relationship("User", back_populates="sessions") user = relationship("User", back_populates="sessions")
def __repr__(self) -> str: def __repr__(self) -> str:
return f"Session(id={self.id}, app_id={self.app_id}, user_id={self.user_id}, location_id={self.location_id}, is_active={self.is_active}, created_at={self.created_at}, h_metadata={self.h_metadata})" return f"Session(id={self.id}, user_id={self.user_id}, location_id={self.location_id}, is_active={self.is_active}, created_at={self.created_at}, h_metadata={self.h_metadata})"
class Message(Base): class Message(Base):
@ -92,7 +92,7 @@ class Message(Base):
h_metadata: Mapped[dict] = mapped_column("metadata", JSONB, default={}) h_metadata: Mapped[dict] = mapped_column("metadata", JSONB, default={})
created_at: Mapped[datetime.datetime] = mapped_column( created_at: Mapped[datetime.datetime] = mapped_column(
DateTime(timezone=True), default=datetime.datetime.utcnow DateTime(timezone=True), default=func.now()
) )
session = relationship("Session", back_populates="messages") session = relationship("Session", back_populates="messages")
metamessages = relationship("Metamessage", back_populates="message") metamessages = relationship("Metamessage", back_populates="message")
@ -112,7 +112,7 @@ class Metamessage(Base):
message = relationship("Message", back_populates="metamessages") message = relationship("Message", back_populates="metamessages")
created_at: Mapped[datetime.datetime] = mapped_column( created_at: Mapped[datetime.datetime] = mapped_column(
DateTime(timezone=True), default=datetime.datetime.utcnow DateTime(timezone=True), default=func.now()
) )
h_metadata: Mapped[dict] = mapped_column("metadata", JSONB, default={}) h_metadata: Mapped[dict] = mapped_column("metadata", JSONB, default={})
@ -127,7 +127,7 @@ class Collection(Base):
) )
name: Mapped[str] = mapped_column(String(512), index=True) name: Mapped[str] = mapped_column(String(512), index=True)
created_at: Mapped[datetime.datetime] = mapped_column( created_at: Mapped[datetime.datetime] = mapped_column(
DateTime(timezone=True), default=datetime.datetime.utcnow DateTime(timezone=True), default=func.now()
) )
h_metadata: Mapped[dict] = mapped_column("metadata", JSONB, default={}) h_metadata: Mapped[dict] = mapped_column("metadata", JSONB, default={})
documents = relationship( documents = relationship(
@ -150,7 +150,7 @@ class Document(Base):
content: Mapped[str] = mapped_column(String(65535)) content: Mapped[str] = mapped_column(String(65535))
embedding = mapped_column(Vector(1536)) embedding = mapped_column(Vector(1536))
created_at: Mapped[datetime.datetime] = mapped_column( created_at: Mapped[datetime.datetime] = mapped_column(
DateTime(timezone=True), default=datetime.datetime.utcnow DateTime(timezone=True), default=func.now()
) )
collection_id = Column(Uuid, ForeignKey("collections.id"), index=True) collection_id = Column(Uuid, ForeignKey("collections.id"), index=True)

View File

@ -1,4 +1,5 @@
import os import os
import traceback
import uuid import uuid
from typing import Optional from typing import Optional
@ -99,6 +100,7 @@ async def create_app(
status_code=406, detail="App with name may already exist" status_code=406, detail="App with name may already exist"
) from e ) from e
except Exception as e: except Exception as e:
print(traceback.format_exc())
raise HTTPException(status_code=400, detail="Unknown Error") from e raise HTTPException(status_code=400, detail="Unknown Error") from e

View File

@ -48,24 +48,6 @@ async def get_collections(
) )
# @router.get("/id/{collection_id}", response_model=schemas.Collection)
# def get_collection_by_id(
# request: Request,
# app_id: uuid.UUID,
# user_id: uuid.UUID,
# collection_id: uuid.UUID,
# db=db,
# ) -> schemas.Collection:
# honcho_collection = 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.get("/name/{name}", response_model=schemas.Collection) @router.get("/name/{name}", response_model=schemas.Collection)
async def get_collection_by_name( async def get_collection_by_name(
request: Request, request: Request,
@ -139,16 +121,11 @@ async def update_collection(
db=db, db=db,
auth=Depends(auth), auth=Depends(auth),
): ):
if collection.name is None: if collection.name is not None and collection.name == "honcho":
raise HTTPException(
status_code=400, detail="invalid request - name cannot be None"
)
if collection.name == "honcho":
raise HTTPException( raise HTTPException(
status_code=406, status_code=406,
detail="error invalid collection configuration - honcho is a reserved name", detail="error invalid collection configuration - honcho is a reserved name",
) )
try: try:
honcho_collection = await crud.update_collection( honcho_collection = await crud.update_collection(
db, db,

View File

@ -11,12 +11,12 @@ from src.dependencies import db
from src.security import auth from src.security import auth
router = APIRouter( router = APIRouter(
prefix="/apps/{app_id}/users/{user_id}/collections/{collection_id}", prefix="/apps/{app_id}/users/{user_id}/collections/{collection_id}/documents",
tags=["documents"], tags=["documents"],
) )
@router.get("/documents", response_model=Page[schemas.Document]) @router.get("", response_model=Page[schemas.Document])
async def get_documents( async def get_documents(
request: Request, request: Request,
app_id: uuid.UUID, app_id: uuid.UUID,
@ -51,7 +51,7 @@ async def get_documents(
@router.get( @router.get(
"/documents/{document_id}", "/{document_id}",
response_model=schemas.Document, response_model=schemas.Document,
) )
async def get_document( async def get_document(
@ -105,7 +105,7 @@ async def query_documents(
) )
@router.post("/documents", response_model=schemas.Document) @router.post("", response_model=schemas.Document)
async def create_document( async def create_document(
request: Request, request: Request,
app_id: uuid.UUID, app_id: uuid.UUID,
@ -130,7 +130,7 @@ async def create_document(
@router.put( @router.put(
"/documents/{document_id}", "/{document_id}",
response_model=schemas.Document, response_model=schemas.Document,
) )
async def update_document( async def update_document(
@ -147,17 +147,22 @@ async def update_document(
raise HTTPException( raise HTTPException(
status_code=400, detail="content and metadata cannot both be None" status_code=400, detail="content and metadata cannot both be None"
) )
return await crud.update_document( try:
db, return await crud.update_document(
document=document, db,
app_id=app_id, document=document,
user_id=user_id, app_id=app_id,
collection_id=collection_id, user_id=user_id,
document_id=document_id, collection_id=collection_id,
) document_id=document_id,
)
except ValueError:
raise HTTPException(
status_code=404, detail="collection not found or does not belong to user"
) from None
@router.delete("/documents/{document_id}") @router.delete("/{document_id}")
async def delete_document( async def delete_document(
request: Request, request: Request,
app_id: uuid.UUID, app_id: uuid.UUID,

View File

@ -33,11 +33,10 @@ async def create_metamessage(
honcho honcho
user_id (str): The User ID representing the user, managed by the user user_id (str): The User ID representing the user, managed by the user
session_id (int): The ID of the Session to add the message to session_id (int): The ID of the Session to add the message to
message (schemas.MessageCreate): The Message object to add containing the metamessage (schemas.MeteamessageCreate): The metamessage creation object
message content and type
Returns: Returns:
schemas.Message: The Message object of the added message schemas.Metamessage: The Metamessage object of the added metamessage
Raises: Raises:
HTTPException: If the session is not found HTTPException: If the session is not found

View File

@ -90,6 +90,9 @@ async def create_session(
return value return value
except ValueError as e: except ValueError as e:
raise HTTPException(status_code=400, detail=str(e)) from e raise HTTPException(status_code=400, detail=str(e)) from e
# except Exception as e:
# print(e)
# raise HTTPException(status_code=400, detail=str(e)) from e
@router.put("/{session_id}", response_model=schemas.Session) @router.put("/{session_id}", response_model=schemas.Session)

View File

@ -1,7 +1,7 @@
import datetime import datetime
import uuid import uuid
from pydantic import BaseModel, Field, validator from pydantic import BaseModel, ConfigDict, Field, field_validator
class AppBase(BaseModel): class AppBase(BaseModel):
@ -25,15 +25,21 @@ class App(AppBase):
metadata: dict metadata: dict
created_at: datetime.datetime created_at: datetime.datetime
@validator("metadata", pre=True, allow_reuse=True) @field_validator("metadata", mode="before")
def fetch_h_metadata(cls, value, values): def fetch_h_metadata(cls, value, info):
if "h_metadata" in values: return info.data.get("h_metadata", {})
return values["h_metadata"] # if "h_metadata" in values:
return {} # return values["h_metadata"]
# return {}
class Config: model_config = ConfigDict(
from_attributes = True from_attributes=True,
json_schema_extra = {"exclude": ["h_metadata"]} json_schema_extra={"exclude": ["h_metadata"]},
)
# class Config:
# from_attributes = True
# json_schema_extra = {"exclude": ["h_metadata"]}
class UserBase(BaseModel): class UserBase(BaseModel):
@ -58,15 +64,22 @@ class User(UserBase):
h_metadata: dict = Field(exclude=True) h_metadata: dict = Field(exclude=True)
metadata: dict metadata: dict
@validator("metadata", pre=True, allow_reuse=True) @field_validator("metadata", mode="before")
def fetch_h_metadata(cls, value, values): def fetch_h_metadata(cls, value, info):
if "h_metadata" in values: return info.data.get("h_metadata", {})
return values["h_metadata"]
return {}
class Config: # @validator("metadata", pre=True, allow_reuse=True)
from_attributes = True # def fetch_h_metadata(cls, value, values):
json_schema_extra = {"exclude": ["h_metadata"]} # if "h_metadata" in values:
# return values["h_metadata"]
# return {}
model_config = ConfigDict(
from_attributes=True,
json_schema_extra={"exclude": ["h_metadata"]},
)
# class Config:
# from_attributes = True
# json_schema_extra = {"exclude": ["h_metadata"]}
class MessageBase(BaseModel): class MessageBase(BaseModel):
@ -92,15 +105,22 @@ class Message(MessageBase):
metadata: dict metadata: dict
created_at: datetime.datetime created_at: datetime.datetime
@validator("metadata", pre=True, allow_reuse=True) @field_validator("metadata", mode="before")
def fetch_h_metadata(cls, value, values): def fetch_h_metadata(cls, value, info):
if "h_metadata" in values: return info.data.get("h_metadata", {})
return values["h_metadata"]
return {}
class Config: # @validator("metadata", pre=True, allow_reuse=True)
from_attributes = True # def fetch_h_metadata(cls, value, values):
json_schema_extra = {"exclude": ["h_metadata"]} # if "h_metadata" in values:
# return values["h_metadata"]
# return {}
model_config = ConfigDict(
from_attributes=True,
json_schema_extra={"exclude": ["h_metadata"]},
)
# class Config:
# from_attributes = True
# json_schema_extra = {"exclude": ["h_metadata"]}
class SessionBase(BaseModel): class SessionBase(BaseModel):
@ -126,15 +146,22 @@ class Session(SessionBase):
metadata: dict metadata: dict
created_at: datetime.datetime created_at: datetime.datetime
@validator("metadata", pre=True, allow_reuse=True) @field_validator("metadata", mode="before")
def fetch_h_metadata(cls, value, values): def fetch_h_metadata(cls, value, info):
if "h_metadata" in values: return info.data.get("h_metadata", {})
return values["h_metadata"]
return {}
class Config: # @validator("metadata", pre=True, allow_reuse=True)
from_attributes = True # def fetch_h_metadata(cls, value, values):
json_schema_extra = {"exclude": ["h_metadata"]} # if "h_metadata" in values:
# return values["h_metadata"]
# return {}
model_config = ConfigDict(
from_attributes=True,
json_schema_extra={"exclude": ["h_metadata"]},
)
# class Config:
# from_attributes = True
# json_schema_extra = {"exclude": ["h_metadata"]}
class MetamessageBase(BaseModel): class MetamessageBase(BaseModel):
@ -163,15 +190,23 @@ class Metamessage(MetamessageBase):
metadata: dict metadata: dict
created_at: datetime.datetime created_at: datetime.datetime
@validator("metadata", pre=True, allow_reuse=True) @field_validator("metadata", mode="before")
def fetch_h_metadata(cls, value, values): def fetch_h_metadata(cls, value, info):
if "h_metadata" in values: return info.data.get("h_metadata", {})
return values["h_metadata"]
return {}
class Config: # @validator("metadata", pre=True, allow_reuse=True)
from_attributes = True # def fetch_h_metadata(cls, value, values):
json_schema_extra = {"exclude": ["h_metadata"]} # if "h_metadata" in values:
# return values["h_metadata"]
# return {}
model_config = ConfigDict(
from_attributes=True,
json_schema_extra={"exclude": ["h_metadata"]},
)
# class Config:
# from_attributes = True
# json_schema_extra = {"exclude": ["h_metadata"]}
class CollectionBase(BaseModel): class CollectionBase(BaseModel):
@ -184,7 +219,7 @@ class CollectionCreate(CollectionBase):
class CollectionUpdate(CollectionBase): class CollectionUpdate(CollectionBase):
name: str name: str | None = None
metadata: dict | None = None metadata: dict | None = None
@ -196,15 +231,23 @@ class Collection(CollectionBase):
metadata: dict metadata: dict
created_at: datetime.datetime created_at: datetime.datetime
@validator("metadata", pre=True, allow_reuse=True) @field_validator("metadata", mode="before")
def fetch_h_metadata(cls, value, values): def fetch_h_metadata(cls, value, info):
if "h_metadata" in values: return info.data.get("h_metadata", {})
return values["h_metadata"]
return {}
class Config: # @validator("metadata", pre=True, allow_reuse=True)
from_attributes = True # def fetch_h_metadata(cls, value, values):
json_schema_extra = {"exclude": ["h_metadata"]} # if "h_metadata" in values:
# return values["h_metadata"]
# return {}
model_config = ConfigDict(
from_attributes=True,
json_schema_extra={"exclude": ["h_metadata"]},
)
# class Config:
# from_attributes = True
# json_schema_extra = {"exclude": ["h_metadata"]}
class DocumentBase(BaseModel): class DocumentBase(BaseModel):
@ -228,15 +271,23 @@ class Document(DocumentBase):
created_at: datetime.datetime created_at: datetime.datetime
collection_id: uuid.UUID collection_id: uuid.UUID
@validator("metadata", pre=True, allow_reuse=True) @field_validator("metadata", mode="before")
def fetch_h_metadata(cls, value, values): def fetch_h_metadata(cls, value, info):
if "h_metadata" in values: return info.data.get("h_metadata", {})
return values["h_metadata"]
return {}
class Config: # @validator("metadata", pre=True, allow_reuse=True)
from_attributes = True # def fetch_h_metadata(cls, value, values):
json_schema_extra = {"exclude": ["h_metadata"]} # if "h_metadata" in values:
# return values["h_metadata"]
# return {}
model_config = ConfigDict(
from_attributes=True,
json_schema_extra={"exclude": ["h_metadata"]},
)
# class Config:
# from_attributes = True
# json_schema_extra = {"exclude": ["h_metadata"]}
class AgentChat(BaseModel): class AgentChat(BaseModel):

View File

@ -17,7 +17,6 @@ async def auth(
credentials: Annotated[HTTPAuthorizationCredentials, Depends(security)], credentials: Annotated[HTTPAuthorizationCredentials, Depends(security)],
): ):
if not USE_AUTH_SERVICE: if not USE_AUTH_SERVICE:
print("Test of Auth")
return True return True
print(credentials) print(credentials)
if not credentials or credentials.credentials != "test": if not credentials or credentials.credentials != "test":

0
tests/__init__.py Normal file
View File

139
tests/conftest.py Normal file
View File

@ -0,0 +1,139 @@
import logging # noqa: I001
import os
import sys
import uuid
import pytest
import pytest_asyncio
from fastapi.testclient import TestClient
from sqlalchemy import text
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine, AsyncSession
from sqlalchemy.engine.url import make_url
from sqlalchemy.exc import OperationalError, ProgrammingError
from sqlalchemy_utils import create_database, database_exists, drop_database
from src import models
from src.db import Base
from src.dependencies import get_db
from src.main import app
logging.basicConfig(
level=logging.DEBUG,
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
stream=sys.stdout, # This ensures the output goes to stdout
)
logger = logging.getLogger(__name__)
# Test database URL
# TODO use environment variable
CONNECTION_URI = make_url(os.getenv("CONNECTION_URI"))
TEST_DB_URL = CONNECTION_URI.set(database="test")
DEFAULT_DB_URL = str(CONNECTION_URI.set(database="postgres"))
def create_test_database(db_url):
"""Helper function create a database if it does not already exist
uses the `sqlalchemy_utils` library to create the database and takes a DB URL
as the input
Args:
db_url (str): Database URL
"""
try:
logger.debug(f"Checking if database exists: {db_url.database}")
if not database_exists(db_url):
logger.info(f"Creating test database: {db_url.database}")
create_database(db_url)
logger.info(f"Test database created successfully: {db_url.database}")
else:
logger.info(f"Database already exists: {db_url.database}")
except Exception as e:
logger.error(f"Error creating database: {e}")
raise
async def setup_test_database(db_url):
"""Helper function to setup the test database
takes a DB URL as input and returns a SQLAlchemy engine
Args:
db_url (str): Database URL
Returns:
engine: SQLAlchemy engine
"""
engine = create_async_engine(str(db_url))
async with engine.connect() as conn:
try:
logger.info("Attempting to create pgvector extension...")
await conn.execute(text("CREATE EXTENSION IF NOT EXISTS vector"))
await conn.commit()
logger.info("pgvector extension created successfully.")
except ProgrammingError as e:
logger.error(f"ProgrammingError: {e}")
raise RuntimeError(
"Failed to create pgvector extension. Make sure it's installed on the PostgreSQL server."
) from e
except OperationalError as e:
logger.error(f"OperationalError: {e}")
raise RuntimeError(
"Failed to connect to the database. Check your connection settings."
) from e
except Exception as e:
logger.error(f"Unexpected error: {e}")
raise
return engine
@pytest_asyncio.fixture(scope="session")
async def db_engine():
create_test_database(TEST_DB_URL)
engine = await setup_test_database(TEST_DB_URL)
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
yield engine
await engine.dispose()
drop_database(TEST_DB_URL)
@pytest_asyncio.fixture(scope="function")
async def db_session(db_engine):
"""Create a database session for the scope of a single test function"""
Session = async_sessionmaker(bind=db_engine, expire_on_commit=False)
async with Session() as session:
yield session
await session.rollback()
@pytest.fixture(scope="function")
def client(db_session):
"""Create a FastAPI TestClient for the scope of a single test function"""
async def override_get_db():
yield db_session
app.dependency_overrides[get_db] = override_get_db
with TestClient(app) as c:
yield c
@pytest_asyncio.fixture(scope="function")
async def sample_data(db_session):
"""Helper function to create test data"""
# Create test app
test_app = models.App(name=str(uuid.uuid4()), metadata={})
db_session.add(test_app)
await db_session.flush()
# Create test user
test_user = models.User(name=str(uuid.uuid4()), app_id=test_app.id, metadata={})
db_session.add(test_user)
await db_session.flush()
yield test_app, test_user
await db_session.rollback()

0
tests/routes/__init__.py Normal file
View File

74
tests/routes/test_apps.py Normal file
View File

@ -0,0 +1,74 @@
import uuid
import pytest
from sqlalchemy.ext.asyncio import AsyncSession
from src import models # Import your SQLAlchemy models
def test_create_app(client):
name = str(uuid.uuid4())
response = client.post("/apps", json={"name": name, "metadata": {"key": "value"}})
print(response)
assert response.status_code == 200
data = response.json()
assert data["name"] == name
assert data["metadata"] == {"key": "value"}
assert "id" in data
def test_get_or_create_app(client):
name = str(uuid.uuid4())
response = client.get(f"/apps/name/{name}")
assert response.status_code == 404
response = client.get(f"/apps/get_or_create/{name}")
assert response.status_code == 200
data = response.json()
assert data["name"] == name
assert "id" in data
def test_get_or_create_existing_app(client):
name = str(uuid.uuid4())
response = client.get(f"/apps/name/{name}")
assert response.status_code == 404
response = client.post("/apps", json={"name": name, "metadata": {"key": "value"}})
assert response.status_code == 200
app1 = response.json()
response = client.get(f"/apps/get_or_create/{name}")
assert response.status_code == 200
app2 = response.json()
assert app1["name"] == app2["name"]
assert app1["id"] == app2["id"]
assert app1["metadata"] == app2["metadata"]
def test_get_app_by_id(client, sample_data):
test_app, _ = sample_data
response = client.get(f"/apps/{test_app.id}")
assert response.status_code == 200
data = response.json()
assert data["name"] == test_app.name
assert data["id"] == str(test_app.id)
def test_get_app_by_name(client, sample_data):
test_app, _ = sample_data
response = client.get(f"/apps/name/{test_app.name}")
assert response.status_code == 200
data = response.json()
assert data["name"] == test_app.name
assert data["id"] == str(test_app.id)
def test_update_app(client, sample_data):
test_app, _ = sample_data
new_name = str(uuid.uuid4())
response = client.put(
f"/apps/{test_app.id}",
json={"name": new_name, "metadata": {"new_key": "new_value"}},
)
assert response.status_code == 200
data = response.json()
assert data["name"] == new_name
assert data["metadata"] == {"new_key": "new_value"}

View File

@ -0,0 +1,95 @@
import uuid
def test_create_collection(client, sample_data) -> None:
test_app, test_user = sample_data
response = client.post(
f"/apps/{test_app.id}/users/{test_user.id}/collections",
json={"name": "test_collection", "metadata": {}},
)
assert response.status_code == 200
data = response.json()
assert data["name"] == "test_collection"
assert data["metadata"] == {}
assert "id" in data
def test_get_collection_by_id(client, sample_data) -> None:
test_app, test_user = sample_data
# Make the collection
response = client.post(
f"/apps/{test_app.id}/users/{test_user.id}/collections",
json={"name": "test_collection", "metadata": {}},
)
assert response.status_code == 200
data = response.json()
# Get the collection
response = client.get(
f"/apps/{test_app.id}/users/{test_user.id}/collections/{data['id']}"
)
assert response.status_code == 200
data = response.json()
assert data["name"] == "test_collection"
assert data["metadata"] == {}
assert "id" in data
def test_get_collection_by_name(client, sample_data) -> None:
test_app, test_user = sample_data
# Make the collection
response = client.post(
f"/apps/{test_app.id}/users/{test_user.id}/collections",
json={"name": "test_collection", "metadata": {}},
)
assert response.status_code == 200
data = response.json()
# Get the collection
response = client.get(
f"/apps/{test_app.id}/users/{test_user.id}/collections/name/test_collection"
)
assert response.status_code == 200
data = response.json()
assert data["name"] == "test_collection"
assert data["metadata"] == {}
assert "id" in data
def test_update_collection(client, sample_data) -> None:
test_app, test_user = sample_data
# Make the collection
response = client.post(
f"/apps/{test_app.id}/users/{test_user.id}/collections",
json={"name": "test_collection", "metadata": {}},
)
assert response.status_code == 200
data = response.json()
# Update the collection
response = client.put(
f"/apps/{test_app.id}/users/{test_user.id}/collections/{data['id']}",
json={"name": "test_collection_updated", "metadata": {"new_key": "new_value"}},
)
assert response.status_code == 200
data = response.json()
assert data["name"] == "test_collection_updated"
assert data["metadata"] == {"new_key": "new_value"}
assert "id" in data
def test_delete_collection(client, sample_data) -> None:
test_app, test_user = sample_data
# Make the collection
response = client.post(
f"/apps/{test_app.id}/users/{test_user.id}/collections",
json={"name": "test_collection", "metadata": {}},
)
assert response.status_code == 200
data = response.json()
# Delete the collection
response = client.delete(
f"/apps/{test_app.id}/users/{test_user.id}/collections/{data['id']}"
)
assert response.status_code == 200
response = client.get(
f"/apps/{test_app.id}/users/{test_user.id}/collections/{data['id']}"
)
assert response.status_code == 404

View File

@ -0,0 +1,96 @@
def test_create_document(client, sample_data):
test_app, test_user = sample_data
# Create a collection
response = client.post(
f"/apps/{test_app.id}/users/{test_user.id}/collections",
json={"name": "test_collection", "metadata": {}},
)
assert response.status_code == 200
data = response.json()
# Create a document
response = client.post(
f"/apps/{test_app.id}/users/{test_user.id}/collections/{data['id']}/documents",
json={"content": "test_text", "metadata": {}},
)
assert response.status_code == 200
data = response.json()
assert data["content"] == "test_text"
assert data["metadata"] == {}
assert "id" in data
def test_get_document(client, sample_data):
test_app, test_user = sample_data
# Create a collection
response = client.post(
f"/apps/{test_app.id}/users/{test_user.id}/collections",
json={"name": "test_collection", "metadata": {}},
)
assert response.status_code == 200
collection = response.json()
# Create a document
response = client.post(
f"/apps/{test_app.id}/users/{test_user.id}/collections/{collection['id']}/documents",
json={"content": "test_text", "metadata": {}},
)
assert response.status_code == 200
document = response.json()
# Get the document
response = client.get(
f"/apps/{test_app.id}/users/{test_user.id}/collections/{collection['id']}/documents/{document['id']}"
)
assert response.status_code == 200
data = response.json()
assert data["content"] == "test_text"
assert data["metadata"] == {}
assert "id" in data
def test_update_document(client, sample_data):
test_app, test_user = sample_data
# Create a collection
response = client.post(
f"/apps/{test_app.id}/users/{test_user.id}/collections",
json={"name": "test_collection", "metadata": {}},
)
assert response.status_code == 200
data = response.json()
# Create a document
response = client.post(
f"/apps/{test_app.id}/users/{test_user.id}/collections/{data['id']}/documents",
json={"content": "test_text", "metadata": {}},
)
assert response.status_code == 200
data = response.json()
# Update the document
response = client.put(
f"/apps/{test_app.id}/users/{test_user.id}/collections/{data['id']}/documents/{data['id']}",
json={"content": "test_text_updated", "metadata": {"new_key": "new_value"}},
)
def test_delete_document(client, sample_data):
test_app, test_user = sample_data
# Create a collection
response = client.post(
f"/apps/{test_app.id}/users/{test_user.id}/collections",
json={"name": "test_collection", "metadata": {}},
)
assert response.status_code == 200
collection = response.json()
# Create a document
response = client.post(
f"/apps/{test_app.id}/users/{test_user.id}/collections/{collection['id']}/documents",
json={"content": "test_text", "metadata": {}},
)
assert response.status_code == 200
document = response.json()
# Delete the document
response = client.delete(
f"/apps/{test_app.id}/users/{test_user.id}/collections/{collection['id']}/documents/{document['id']}"
)
assert response.status_code == 200
response = client.get(
f"/apps/{test_app.id}/users/{test_user.id}/collections/{collection['id']}/documents/{document['id']}"
)
assert response.status_code == 404

View File

@ -0,0 +1,80 @@
import pytest
from src import models # Import your SQLAlchemy models
@pytest.mark.asyncio
async def test_create_message(client, db_session, sample_data):
test_app, test_user = sample_data
# Create a test session
test_session = models.Session(
user_id=test_user.id, location_id="test_location", metadata={}
)
db_session.add(test_session)
await db_session.commit()
response = client.post(
f"/apps/{test_app.id}/users/{test_user.id}/sessions/{test_session.id}/messages",
json={
"content": "Test message",
"is_user": True,
"metadata": {"message_key": "message_value"},
},
)
assert response.status_code == 200
data = response.json()
assert data["content"] == "Test message"
assert data["is_user"] is True
assert data["metadata"] == {"message_key": "message_value"}
assert "id" in data
@pytest.mark.asyncio
async def test_get_messages(client, db_session, sample_data):
test_app, test_user = sample_data
# Create a test session and message
test_session = models.Session(
user_id=test_user.id, location_id="test_location", metadata={}
)
db_session.add(test_session)
await db_session.commit()
test_message = models.Message(
session_id=test_session.id, content="Test message", is_user=True, metadata={}
)
db_session.add(test_message)
await db_session.commit()
response = client.get(
f"/apps/{test_app.id}/users/{test_user.id}/sessions/{test_session.id}/messages"
)
assert response.status_code == 200
data = response.json()
assert "items" in data
assert len(data["items"]) > 0
assert data["items"][0]["content"] == "Test message"
assert data["items"][0]["is_user"] is True
assert data["items"][0]["metadata"] == {}
@pytest.mark.asyncio
async def test_update_message(client, db_session, sample_data):
test_app, test_user = sample_data
# Create a test session and message
test_session = models.Session(
user_id=test_user.id, location_id="test_location", metadata={}
)
db_session.add(test_session)
await db_session.commit()
test_message = models.Message(
session_id=test_session.id, content="Test message", is_user=True, metadata={}
)
db_session.add(test_message)
await db_session.commit()
response = client.put(
f"/apps/{test_app.id}/users/{test_user.id}/sessions/{test_session.id}/messages/{test_message.id}",
json={"metadata": {"new_key": "new_value"}},
)
assert response.status_code == 200
data = response.json()
assert data["metadata"] == {"new_key": "new_value"}

View File

@ -0,0 +1,101 @@
import pytest
from src import models # Import your SQLAlchemy models
@pytest.mark.asyncio
async def test_create_metamessage(client, db_session, sample_data):
test_app, test_user = sample_data
# Create a test session
test_session = models.Session(
user_id=test_user.id, location_id="test_location", metadata={}
)
db_session.add(test_session)
await db_session.commit()
test_message = models.Message(
session_id=test_session.id, content="Test message", is_user=True, metadata={}
)
db_session.add(test_message)
await db_session.commit()
response = client.post(
f"/apps/{test_app.id}/users/{test_user.id}/sessions/{test_session.id}/metamessages",
json={
"message_id": str(test_message.id),
"content": "Test Metamessage",
"metadata": {},
"metamessage_type": "test_type",
},
)
assert response.status_code == 200
data = response.json()
assert data["message_id"] == str(test_message.id)
assert data["content"] == "Test Metamessage"
assert data["metadata"] == {}
assert data["metamessage_type"] == "test_type"
@pytest.mark.asyncio
async def test_get_metamessage(client, db_session, sample_data):
test_app, test_user = sample_data
# Create a test session
test_session = models.Session(
user_id=test_user.id, location_id="test_location", metadata={}
)
db_session.add(test_session)
await db_session.commit()
test_message = models.Message(
session_id=test_session.id, content="Test message", is_user=True, metadata={}
)
db_session.add(test_message)
await db_session.commit()
test_metamessage = models.Metamessage(
message_id=test_message.id,
content="Test Metamessage",
metadata={},
metamessage_type="test_type",
)
db_session.add(test_metamessage)
await db_session.commit()
response = client.get(
f"/apps/{test_app.id}/users/{test_user.id}/sessions/{test_session.id}/metamessages/{test_metamessage.id}?message_id={test_message.id}"
)
assert response.status_code == 200
data = response.json()
assert data["message_id"] == str(test_message.id)
assert data["content"] == "Test Metamessage"
assert data["metadata"] == {}
assert data["metamessage_type"] == "test_type"
@pytest.mark.asyncio
async def test_update_metamessage(client, db_session, sample_data):
test_app, test_user = sample_data
# Create a test session
test_session = models.Session(
user_id=test_user.id, location_id="test_location", metadata={}
)
db_session.add(test_session)
await db_session.commit()
test_message = models.Message(
session_id=test_session.id, content="Test message", is_user=True, metadata={}
)
db_session.add(test_message)
await db_session.commit()
test_metamessage = models.Metamessage(
message_id=test_message.id,
content="Test Metamessage",
metadata={},
metamessage_type="test_type",
)
db_session.add(test_metamessage)
await db_session.commit()
response = client.put(
f"/apps/{test_app.id}/users/{test_user.id}/sessions/{test_session.id}/metamessages/{test_metamessage.id}",
json={"message_id": str(test_message.id), "metadata": {"new_key": "new_value"}},
)
assert response.status_code == 200
data = response.json()
assert data["metadata"] == {"new_key": "new_value"}

View File

@ -0,0 +1,76 @@
import pytest
from src import models # Import your SQLAlchemy models
def test_create_session(client, sample_data):
test_app, test_user = sample_data
response = client.post(
f"/apps/{test_app.id}/users/{test_user.id}/sessions",
json={
"location_id": "test_location",
"metadata": {"session_key": "session_value"},
},
)
assert response.status_code == 200
data = response.json()
assert data["location_id"] == "test_location"
assert data["metadata"] == {"session_key": "session_value"}
assert "id" in data
@pytest.mark.asyncio
async def test_get_sessions(client, db_session, sample_data):
test_app, test_user = sample_data
# Create a test session
test_session = models.Session(
user_id=test_user.id, location_id="test_location", metadata={}
)
db_session.add(test_session)
await db_session.commit()
response = client.get(f"/apps/{test_app.id}/users/{test_user.id}/sessions")
assert response.status_code == 200
data = response.json()
assert "items" in data
assert len(data["items"]) > 0
assert data["items"][0]["location_id"] == "test_location"
@pytest.mark.asyncio
async def test_update_session(client, db_session, sample_data):
test_app, test_user = sample_data
# Create a test session
test_session = models.Session(
user_id=test_user.id, location_id="test_location", metadata={}
)
db_session.add(test_session)
await db_session.commit()
response = client.put(
f"/apps/{test_app.id}/users/{test_user.id}/sessions/{test_session.id}",
json={"metadata": {"new_key": "new_value"}},
)
assert response.status_code == 200
data = response.json()
assert data["metadata"] == {"new_key": "new_value"}
@pytest.mark.asyncio
async def test_delete_session(client, db_session, sample_data):
test_app, test_user = sample_data
# Create a test session
test_session = models.Session(
user_id=test_user.id, location_id="test_location", metadata={}
)
db_session.add(test_session)
await db_session.commit()
response = client.delete(
f"/apps/{test_app.id}/users/{test_user.id}/sessions/{test_session.id}"
)
assert response.status_code == 200
response = client.get(
f"/apps/{test_app.id}/users/{test_user.id}/sessions/{test_session.id}"
)
data = response.json()
assert data["is_active"] is False

View File

@ -0,0 +1,69 @@
import uuid
def test_create_user(client, sample_data):
test_app, _ = sample_data
name = str(uuid.uuid4())
response = client.post(
f"/apps/{test_app.id}/users",
json={"name": name, "metadata": {"user_key": "user_value"}},
)
assert response.status_code == 200
data = response.json()
assert data["name"] == name
assert data["metadata"] == {"user_key": "user_value"}
assert "id" in data
def test_get_user_by_id(client, sample_data):
test_app, test_user = sample_data
response = client.get(f"/apps/{test_app.id}/users/{test_user.id}")
assert response.status_code == 200
data = response.json()
assert data["name"] == test_user.name
assert data["id"] == str(test_user.id)
def test_get_user_by_name(client, sample_data):
test_app, test_user = sample_data
response = client.get(f"/apps/{test_app.id}/users/name/{test_user.name}")
assert response.status_code == 200
data = response.json()
assert data["name"] == test_user.name
assert data["id"] == str(test_user.id)
def test_get_or_create_user(client, sample_data):
test_app, _ = sample_data
name = str(uuid.uuid4())
response = client.get(f"/apps/{test_app.id}/users/name/{name}")
assert response.status_code == 404
response = client.get(f"/apps/{test_app.id}/users/get_or_create/{name}")
assert response.status_code == 200
data = response.json()
assert data["name"] == name
assert "id" in data
# def test_get_users(client, sample_data):
# test_app, _ = sample_data
# response = client.get(f"/apps/{test_app.id}/users")
# assert response.status_code == 200
# data = response.json()
# assert "items" in data
# assert len(data["items"]) > 0
def test_update_user(client, sample_data):
test_app, test_user = sample_data
new_name = str(uuid.uuid4())
response = client.put(
f"/apps/{test_app.id}/users/{test_user.id}",
json={"name": new_name, "metadata": {"new_key": "new_value"}},
)
assert response.status_code == 200
data = response.json()
print(new_name)
print(data)
assert data["name"] == new_name
assert data["metadata"] == {"new_key": "new_value"}