Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -76,7 +76,7 @@ cd ta-ai
cd frontend
npm install

# Install backend dependencies
# Install backend dependencies (use Python 3.11 for local dev)
cd ../backend
pip install -r requirements.txt

Expand All @@ -86,6 +86,13 @@ terraform init
terraform plan
```

### Observability
- JSON logs and request timing are enabled by default.
- Optional tracing: set `ENABLE_OTEL=1`. Configure `OTEL_EXPORTER_OTLP_ENDPOINT` and `OTEL_SERVICE_NAME` as needed.

### API Reference
See `docs/api-reference.md` for endpoints and environment variables.

## Cost Estimation

| Component | Monthly Cost |
Expand Down
3 changes: 2 additions & 1 deletion backend/requirements.txt
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,8 @@ python-dotenv==1.0.0
openai==1.6.1

# Database
psycopg2-binary==2.9.9
psycopg2-binary==2.9.9 ; python_version < "3.13"
psycopg[binary]==3.2.1 ; python_version >= "3.13"
pgvector==0.2.4
sqlalchemy==2.0.23
alembic==1.13.1
Expand Down
3 changes: 2 additions & 1 deletion backend/src/db.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,8 @@
# Load environment variables from .env
load_dotenv()

DATABASE_URL = os.getenv("DATABASE_URL", "postgresql:///ta_ai")
# Default to SQLite for local development to avoid requiring Postgres drivers
DATABASE_URL = os.getenv("DATABASE_URL", "sqlite:///./ta_ai.db")

# SQLAlchemy engine and session
engine = create_engine(
Expand Down
19 changes: 14 additions & 5 deletions backend/src/functions/ingest/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,11 +5,17 @@
from services.document_parser import parse_document
from db import SessionLocal
from models.models import Chunk
import openai
from openai import OpenAI
import tiktoken

# Load OpenAI API key from env
openai.api_key = os.getenv("AZURE_OPENAI_API_KEY") or os.getenv("OPENAI_API_KEY")
# Lazily initialized OpenAI client. Tests may monkeypatch this symbol.
client: OpenAI | None = None

def get_client() -> OpenAI:
global client
if client is None:
client = OpenAI()
return client

# Configure embedding model
EMBEDDING_MODEL = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small")
Expand Down Expand Up @@ -51,8 +57,11 @@ def main(req: func.HttpRequest) -> func.HttpResponse:
try:
for chunk in chunk_text(text, MAX_TOKENS):
# generate embedding
resp = openai.Embedding.create(input=chunk, model=EMBEDDING_MODEL)
embedding = resp["data"][0]["embedding"]
if os.getenv("MOCK_OPENAI") == "1":
embedding = [0.1, 0.2, 0.3]
else:
resp = get_client().embeddings.create(input=chunk, model=EMBEDDING_MODEL)
embedding = resp.data[0].embedding
# persist chunk
db_chunk = Chunk(course_id=course_id, text=chunk, embedding=embedding)
session.add(db_chunk)
Expand Down
125 changes: 122 additions & 3 deletions backend/src/main.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,20 @@
"""
Main FastAPI application for TA AI backend
"""
from fastapi import FastAPI
from fastapi import FastAPI, Request
from fastapi.middleware.cors import CORSMiddleware
import os
import time
import uuid
import json
import logging
from datetime import datetime, timezone
from dotenv import load_dotenv
from .db import engine, Base
# Support running both as package (uvicorn src.main:app) and as script (python src/main.py)
try:
from db import engine, Base # when src/ is on sys.path
except ImportError: # pragma: no cover
from src.db import engine, Base # fallback when launched as a package

# Load environment variables
load_dotenv()
Expand All @@ -31,13 +40,120 @@
allow_headers=["*"],
)

class JsonFormatter(logging.Formatter):
def format(self, record: logging.LogRecord) -> str:
payload = {
"time": datetime.now(timezone.utc).isoformat(),
"level": record.levelname,
"logger": record.name,
"message": record.getMessage(),
}
for key in (
"request_id",
"request_path",
"request_method",
"status_code",
"duration_ms",
"client_host",
"user_agent",
):
val = getattr(record, key, None)
if val is not None:
payload[key] = val
return json.dumps(payload, ensure_ascii=False)


def setup_json_logging() -> None:
handler = logging.StreamHandler()
handler.setFormatter(JsonFormatter())
root = logging.getLogger()
root.handlers.clear()
root.addHandler(handler)
root.setLevel(logging.INFO)


def setup_opentelemetry_if_enabled() -> None:
if os.getenv("ENABLE_OTEL") != "1":
return
try:
# Lazy import; optional dependency
from opentelemetry import trace
from opentelemetry.sdk.resources import Resource
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.sdk.trace.export import BatchSpanProcessor, ConsoleSpanExporter
try:
from opentelemetry.exporter.otlp.proto.http.trace_exporter import OTLPSpanExporter
otlp_endpoint = os.getenv("OTEL_EXPORTER_OTLP_ENDPOINT")
exporter = OTLPSpanExporter(endpoint=otlp_endpoint) if otlp_endpoint else ConsoleSpanExporter()
except Exception:
exporter = ConsoleSpanExporter()

resource = Resource.create({"service.name": os.getenv("OTEL_SERVICE_NAME", "ta-ai-backend")})
provider = TracerProvider(resource=resource)
provider.add_span_processor(BatchSpanProcessor(exporter))
trace.set_tracer_provider(provider)
logging.getLogger("otel").info("OpenTelemetry tracing enabled")
except Exception as e:
logging.getLogger("otel").warning(f"OpenTelemetry not enabled: {e}")

@app.on_event("startup")
async def on_startup():
# Import models to register metadata
import src.models.models # noqa: F401
# Create database tables
print("[Startup] Creating database tables...")
Base.metadata.create_all(bind=engine)
setup_json_logging()
setup_opentelemetry_if_enabled()


@app.middleware("http")
async def access_log_middleware(request: Request, call_next):
start = time.perf_counter()
request_id = str(uuid.uuid4())
status = 500
span_ctx = None
tracer = None
try:
# Optional span
try:
from opentelemetry import trace
tracer = trace.get_tracer(__name__)
span_ctx = tracer.start_as_current_span("http.request")
span_ctx.__enter__()
except Exception:
tracer = None
try:
response = await call_next(request)
status = response.status_code
return response
finally:
if span_ctx is not None:
try:
span = trace.get_current_span() # type: ignore[name-defined]
if span is not None:
span.set_attribute("http.method", request.method)
span.set_attribute("http.target", request.url.path)
span.set_attribute("http.status_code", status)
except Exception:
pass
try:
span_ctx.__exit__(None, None, None)
except Exception:
pass
duration_ms = int((time.perf_counter() - start) * 1000)
logging.getLogger("access").info(
"request",
extra={
"request_id": request_id,
"request_path": request.url.path,
"request_method": request.method,
"status_code": locals().get("status", 500),
"duration_ms": duration_ms,
"client_host": request.client.host if request.client else None,
"user_agent": request.headers.get("user-agent"),
},
)

@app.get("/api/health")
async def health_check():
Expand All @@ -54,7 +170,10 @@ async def test_endpoint():

# QA query endpoint
from pydantic import BaseModel
from src.services.qa_service import generate_answer
try:
from services.qa_service import generate_answer
except ImportError: # pragma: no cover
from src.services.qa_service import generate_answer

class QueryRequest(BaseModel):
course_id: int
Expand Down
9 changes: 7 additions & 2 deletions backend/src/models/models.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,12 @@
from sqlalchemy import Column, Integer, String, ForeignKey, DateTime, Text
from sqlalchemy.dialects.postgresql import ARRAY
from pgvector.sqlalchemy import Vector
from sqlalchemy import JSON
from datetime import datetime
from src.db import Base
try:
from db import Base
except ImportError: # pragma: no cover
from src.db import Base

class Course(Base):
__tablename__ = "courses"
Expand All @@ -29,7 +33,8 @@ class QuestionLog(Base):
user_id = Column(Integer, ForeignKey("users.id"), index=True, nullable=False)
question = Column(Text, nullable=False)
answer = Column(Text, nullable=False)
citations = Column(ARRAY(Integer), nullable=True)
# Use ARRAY on Postgres, fallback to JSON when using SQLite for local dev
citations = Column(JSON, nullable=True)
timestamp = Column(DateTime, default=datetime.utcnow, nullable=False)

class Feedback(Base):
Expand Down
84 changes: 84 additions & 0 deletions backend/src/services/embedding_service.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,84 @@
import os
from openai import OpenAI
from sqlalchemy import text
try:
from db import SessionLocal
except ImportError: # pragma: no cover
from src.db import SessionLocal

# Lazily initialized OpenAI client. Tests may monkeypatch this symbol.
client: OpenAI | None = None


def get_client() -> OpenAI:
global client
if client is None:
client = OpenAI()
return client

# Embedding model
EMBEDDING_MODEL = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small")
# Number of results
DEFAULT_K = int(os.getenv("KNN_K", 5))


def embed_query(query: str) -> list[float]:
"""Generate embedding for the input query using OpenAI.

When MOCK_OPENAI=1, returns a deterministic small vector to avoid external calls.
"""
if os.getenv("MOCK_OPENAI") == "1":
return [0.1, 0.2, 0.3]

response = get_client().embeddings.create(input=query, model=EMBEDDING_MODEL)
return response.data[0].embedding


def retrieve_chunks(course_id: int, query_embedding: list[float], k: int = DEFAULT_K) -> list[dict]:
"""Retrieve the top-k similar chunks.

- Uses pgvector cosine distance when connected to Postgres.
- Falls back to a simple top-k by ID for SQLite or when MOCK_OPENAI is enabled.
"""
session = SessionLocal()
try:
use_simple = os.getenv("MOCK_OPENAI") == "1"
try:
dialect_name = session.bind.dialect.name # type: ignore[attr-defined]
if dialect_name == "sqlite":
use_simple = True
except Exception:
pass

if use_simple:
# Simple fallback: return up to k chunks for course_id with dummy distance
try:
# Lazy import to avoid circular dependency
from models.models import Chunk # type: ignore
except Exception:
from src.models.models import Chunk # type: ignore
results = (
session.query(Chunk)
.filter(Chunk.course_id == course_id)
.limit(k)
.all()
)
return [{"id": c.id, "text": c.text, "distance": 0.0} for c in results]

# Postgres with pgvector path
sql = text(
"SELECT id, text, embedding <-> :query_embedding AS distance "
"FROM chunks "
"WHERE course_id = :course_id "
"ORDER BY distance "
"LIMIT :k"
)
result = session.execute(
sql,
{"query_embedding": query_embedding, "course_id": course_id, "k": k},
)
mapped = result.mappings()
rows = mapped.all() if hasattr(mapped, "all") else list(mapped)
return rows
finally:
session.close()
36 changes: 23 additions & 13 deletions backend/src/services/qa_service.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,20 @@
import os
import openai
from openai import OpenAI
from typing import List, Dict
from services.embedding_service import embed_query, retrieve_chunks
try:
from services.embedding_service import embed_query, retrieve_chunks
except ImportError: # pragma: no cover
from src.services.embedding_service import embed_query, retrieve_chunks

# Load Azure OpenAI configuration
openai.api_key = os.getenv("AZURE_OPENAI_API_KEY") or os.getenv("OPENAI_API_KEY")
openai.api_type = os.getenv("OPENAI_API_TYPE", "azure")
openai.api_base = os.getenv("AZURE_OPENAI_ENDPOINT", os.getenv("OPENAI_API_BASE"))
openai.api_version = os.getenv("AZURE_OPENAI_API_VERSION", os.getenv("OPENAI_API_VERSION"))
# Lazily initialized OpenAI client
client: OpenAI | None = None

def get_client() -> OpenAI:
global client
if client is None:
client = OpenAI()
return client

# Models and defaults
CHAT_MODEL = os.getenv("CHAT_MODEL", "gpt-4-turbo")
Expand Down Expand Up @@ -39,13 +46,16 @@ def generate_answer(question: str, course_id: int, k: int = DEFAULT_K) -> Dict:
{"role": "user", "content": f"Here are relevant snippets:\n{context}\n\nQuestion: {question}"},
]

# 4. Call ChatCompletion
response = openai.ChatCompletion.create(
model=CHAT_MODEL,
messages=messages,
temperature=0.2,
)
answer = response["choices"][0]["message"]["content"]
# 4. Call Chat Completions API (responses API if streaming desired)
if os.getenv("MOCK_OPENAI") == "1":
answer = "This is a mocked answer. [ID 1][ID 2]"
else:
response = get_client().chat.completions.create(
model=CHAT_MODEL,
messages=messages,
temperature=0.2,
)
answer = response.choices[0].message.content

# 5. Return answer and citations
citations: List[int] = [row["id"] for row in chunks]
Expand Down
Loading
Loading