Skip to content
Merged
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: 7 additions & 2 deletions Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,8 @@ PYTEST_FLAGS ?= -q
TEST_DB_HOST ?= localhost
TEST_DB_PORT ?= 5432
TEST_DB_NAME ?= efficientai
# Dedicated DB for test-docker-db / local Postgres runs (matches CI; not the dev DB).
TEST_DOCKER_DB_NAME ?= efficientai_test
TEST_DB_USER ?= efficientai
TEST_DB_PASSWORD ?= password
TEST_DATABASE_URL ?= postgresql://$(TEST_DB_USER):$(TEST_DB_PASSWORD)@$(TEST_DB_HOST):$(TEST_DB_PORT)/$(TEST_DB_NAME)
Expand Down Expand Up @@ -49,8 +51,11 @@ test-parallel: check-pytest-xdist ## Run the full test suite in parallel (pytest
$(PYTEST) tests $(PYTEST_FLAGS) -n auto --dist loadscope $(PYTEST_ARGS)

test-docker-db: check-pytest ## Run tests against running Docker Compose Postgres
TEST_DATABASE_URL="$(TEST_DATABASE_URL)" DATABASE_URL="$(TEST_DATABASE_URL)" \
POSTGRES_HOST="$(TEST_DB_HOST)" POSTGRES_PORT="$(TEST_DB_PORT)" POSTGRES_DB="$(TEST_DB_NAME)" \
@PGPASSWORD="$(TEST_DB_PASSWORD)" psql -h "$(TEST_DB_HOST)" -p "$(TEST_DB_PORT)" -U "$(TEST_DB_USER)" -d postgres -tc "SELECT 1 FROM pg_database WHERE datname='$(TEST_DOCKER_DB_NAME)'" | grep -q 1 \
|| PGPASSWORD="$(TEST_DB_PASSWORD)" psql -h "$(TEST_DB_HOST)" -p "$(TEST_DB_PORT)" -U "$(TEST_DB_USER)" -d postgres -c "CREATE DATABASE \"$(TEST_DOCKER_DB_NAME)\";"
TEST_DATABASE_URL="postgresql://$(TEST_DB_USER):$(TEST_DB_PASSWORD)@$(TEST_DB_HOST):$(TEST_DB_PORT)/$(TEST_DOCKER_DB_NAME)" \
DATABASE_URL="postgresql://$(TEST_DB_USER):$(TEST_DB_PASSWORD)@$(TEST_DB_HOST):$(TEST_DB_PORT)/$(TEST_DOCKER_DB_NAME)" \
POSTGRES_HOST="$(TEST_DB_HOST)" POSTGRES_PORT="$(TEST_DB_PORT)" POSTGRES_DB="$(TEST_DOCKER_DB_NAME)" \
POSTGRES_USER="$(TEST_DB_USER)" POSTGRES_PASSWORD="$(TEST_DB_PASSWORD)" \
$(PYTEST) tests $(PYTEST_FLAGS) $(PYTEST_ARGS)

Expand Down
152 changes: 139 additions & 13 deletions app/api/v1/routes/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,13 @@
consume_reference_code,
validate_reference_code_for_signup,
)
from app.services.invitation_service import (
InvitationError,
accept_invitation as accept_invitation_record,
get_invitation_preview,
get_valid_pending_invitation_by_token,
)
from app.api.v1.routes.profile import get_current_user

router = APIRouter(prefix="/auth", tags=["Authentication"])

Expand Down Expand Up @@ -95,12 +102,28 @@ class SignupRequest(BaseModel):
first_name: Optional[str] = Field(default=None, max_length=255)
last_name: Optional[str] = Field(default=None, max_length=255)
reference_code: Optional[str] = Field(default=None, max_length=64)
invite_token: Optional[str] = Field(default=None, max_length=255)


class InvitationPreviewResponse(BaseModel):
organization_name: Optional[str] = None
email: str
role: str
expires_at: datetime
status: str
user_exists: bool = False
has_password: bool = False


class AcceptInviteByTokenRequest(BaseModel):
token: str = Field(min_length=1, max_length=255)


class LoginRequest(BaseModel):
email: EmailStr
password: str
organization_id: Optional[str] = None
invite_token: Optional[str] = Field(default=None, max_length=255)


class LoginOrgOption(BaseModel):
Expand Down Expand Up @@ -290,6 +313,38 @@ def _issue_session_tokens(
)


@router.get("/invitations/preview/{token}", response_model=InvitationPreviewResponse)
def preview_invitation(token: str, db: Session = Depends(get_db)) -> InvitationPreviewResponse:
"""Public preview of an organization invite (no auth required)."""
try:
preview = get_invitation_preview(db, token)
except InvitationError as exc:
raise HTTPException(status_code=exc.status_code, detail=exc.detail) from exc
return InvitationPreviewResponse(**preview)


@router.post("/invitations/accept-by-token", response_model=TokenResponse)
def accept_invitation_by_token(
payload: AcceptInviteByTokenRequest,
current_user: User = Depends(get_current_user),
db: Session = Depends(get_db),
) -> TokenResponse:
"""Accept an invitation and return a session scoped to the invited organization."""
try:
invitation = get_valid_pending_invitation_by_token(db, payload.token)
member = accept_invitation_record(db, invitation, current_user)
except InvitationError as exc:
raise HTTPException(status_code=exc.status_code, detail=exc.detail) from exc

role_value = member.role.value if hasattr(member.role, "value") else member.role
return _issue_session_tokens(
db,
user=current_user,
organization_id=invitation.organization_id,
role_value=role_value,
)


@router.post("/signup", response_model=TokenResponse)
def signup(payload: SignupRequest, db: Session = Depends(get_db)) -> TokenResponse:
"""
Expand All @@ -307,31 +362,81 @@ def signup(payload: SignupRequest, db: Session = Depends(get_db)) -> TokenRespon
)

reference_row = None
if settings.AUTH_GATED_SIGNUP_ENABLED:
invite_token = (payload.invite_token or "").strip() or None
pending_invitation = None

if invite_token:
try:
pending_invitation = get_valid_pending_invitation_by_token(db, invite_token)
except InvitationError as exc:
raise HTTPException(status_code=exc.status_code, detail=exc.detail) from exc
if pending_invitation.email.lower() != payload.email.lower():
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="Email does not match the invitation.",
)
elif settings.AUTH_GATED_SIGNUP_ENABLED:
reference_row = validate_reference_code_for_signup(db, payload.reference_code)

existing = db.query(User).filter(User.email == payload.email).first()
if existing:
if existing and existing.password_hash:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail="An account with this email already exists. Try signing in instead.",
)

org_name = (payload.organization_name or payload.email.split("@")[0] + "'s Org").strip()
_validate_password_or_400(payload.password)

user = User(
email=payload.email,
password_hash=hash_password(payload.password),
first_name=payload.first_name,
last_name=payload.last_name,
name=((payload.first_name or "") + " " + (payload.last_name or "")).strip() or None,
is_active=True,
auth_provider="local",
)
if existing and not existing.password_hash:
user = existing
user.password_hash = hash_password(payload.password)
if payload.first_name:
user.first_name = payload.first_name
if payload.last_name:
user.last_name = payload.last_name
if payload.first_name or payload.last_name:
user.name = ((payload.first_name or "") + " " + (payload.last_name or "")).strip() or user.name
user.is_active = True
if not user.auth_provider:
user.auth_provider = "local"
db.flush()
else:
user = User(
email=payload.email,
password_hash=hash_password(payload.password),
first_name=payload.first_name,
last_name=payload.last_name,
name=((payload.first_name or "") + " " + (payload.last_name or "")).strip() or None,
is_active=True,
auth_provider="local",
)
db.add(user)
db.flush()

if pending_invitation is not None:
member = accept_invitation_record(
db,
pending_invitation,
user,
require_email_match=True,
)
if reference_row is not None:
consume_reference_code(db, reference_row)
user.last_login_at = datetime.now(timezone.utc)
db.commit()
db.refresh(user)

role_value = member.role.value if hasattr(member.role, "value") else member.role
return _issue_session_tokens(
db,
user=user,
organization_id=pending_invitation.organization_id,
role_value=role_value,
)

org_name = (payload.organization_name or payload.email.split("@")[0] + "'s Org").strip()
organization = Organization(name=org_name)
db.add(organization)
db.add(user)
db.flush()

membership = OrganizationMember(
Expand Down Expand Up @@ -395,6 +500,27 @@ def login(payload: LoginRequest, db: Session = Depends(get_db)) -> LoginResponse
.order_by(OrganizationMember.joined_at.asc())
.all()
)

invite_token = (payload.invite_token or "").strip() or None
if not memberships and invite_token:
try:
invitation = get_valid_pending_invitation_by_token(db, invite_token)
if invitation.email.lower() != user.email.lower():
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="This invitation was sent to a different email address.",
)
member = accept_invitation_record(db, invitation, user)
organization = (
db.query(Organization)
.filter(Organization.id == invitation.organization_id)
.first()
)
if organization is not None and organization.is_active:
memberships = [(member, organization)]
except InvitationError as exc:
raise HTTPException(status_code=exc.status_code, detail=exc.detail) from exc

if not memberships:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
Expand Down
44 changes: 41 additions & 3 deletions app/api/v1/routes/call_import_evaluations.py
Original file line number Diff line number Diff line change
Expand Up @@ -563,6 +563,29 @@ def _serialize_eval(
user_emails = emails_for_user_ids(db, user_ids_from_evaluations([row]))
created_email, updated_email = actor_emails_for_evaluation(row, user_emails)

from app.workers.tasks.evaluate_call_import_row_core import (
count_distinct_llm_configs_for_metrics,
)

selected_set = set(selected_ids)
scoring_metrics = [m for m in metrics if m.id in selected_set]
expected_llm_calls_per_row = count_distinct_llm_configs_for_metrics(
scoring_metrics,
overrides=(
row.metric_llm_overrides
if isinstance(row.metric_llm_overrides, dict)
else {}
),
run_provider=(row.llm_provider or "").strip() or None,
run_model=(row.llm_model or "").strip() or None,
run_llm_config=(
row.llm_config if isinstance(getattr(row, "llm_config", None), dict) else None
),
run_credential_id=(
str(row.llm_credential_id) if row.llm_credential_id else None
),
)

return CallImportEvaluationResponse(
id=row.id,
call_import_id=row.call_import_id,
Expand Down Expand Up @@ -618,6 +641,7 @@ def _serialize_eval(
),
transcript_source=(row.transcript_source or "diarised"),
sibling_evaluation_ids=list(sibling_evaluation_ids or []),
expected_llm_calls_per_row=expected_llm_calls_per_row,
started_at=row.started_at,
finished_at=row.finished_at,
created_at=row.created_at,
Expand Down Expand Up @@ -1014,7 +1038,6 @@ async def create_call_import_evaluation(
from app.models.enums import CallImportParameterType, CallImportStatus
from app.services.call_imports.bulk_ops import (
count_all_source_rows,
count_completed_source_rows,
count_source_rows_with_production_transcript,
)

Expand Down Expand Up @@ -1124,15 +1147,30 @@ async def create_call_import_evaluation(
db.refresh(call_import)
starting_from_mapped = True

all_source_row_count = count_all_source_rows(db, call_import.id)

if use_diarised:
total_row_count = count_completed_source_rows(db, call_import.id)
# Diarized runs materialize every import row; recording fetch may
# still be pending when switching from a production eval-primary run.
total_row_count = all_source_row_count
else:
# Production runs score CSV text — rows need not wait for
# recording fetch to finish before they are evaluable.
total_row_count = count_source_rows_with_production_transcript(
db, call_import.id
)

if (
use_diarised
and not starting_from_mapped
and all_source_row_count > 0
and call_import.status != CallImportStatus.DELETING
):
call_import.status = CallImportStatus.PROCESSING
stamp_call_import_actor(call_import, principal)
db.commit()
db.refresh(call_import)

requested_sources: List[str] = list(payload.transcript_sources)

if (
Expand Down Expand Up @@ -1209,7 +1247,7 @@ def _name_for_source(source: str) -> Optional[str]:
primary_evaluation = created_evaluations[0]
sibling_ids = [e.id for e in created_evaluations[1:]]

if not total_row_count and not starting_from_mapped:
if not all_source_row_count and not starting_from_mapped:
for evaluation in created_evaluations:
evaluation.status = "completed"
db.commit()
Expand Down
8 changes: 7 additions & 1 deletion app/api/v1/routes/call_imports.py
Original file line number Diff line number Diff line change
Expand Up @@ -148,7 +148,10 @@ def _serialize_call_import(
user_emails: Optional[Dict[UUID, str]] = None,
) -> CallImportResponse:
"""Catalog parent fields; counters come from SQL rollup (not Redis merge)."""
from app.services.call_imports.bulk_ops import rollup_call_import_batch_status
from app.services.call_imports.bulk_ops import (
_latest_evaluation_status,
rollup_call_import_batch_status,
)
from app.services.call_imports.progress_counters import (
clear_import_progress_redis,
read_import_progress,
Expand Down Expand Up @@ -185,6 +188,9 @@ def _serialize_call_import(
update={
"completed_rows": completed,
"failed_rows": failed,
"latest_evaluation_status": _latest_evaluation_status(
db, call_import.id
),
"created_by_email": created_email,
"last_updated_by_email": updated_email,
}
Expand Down
Loading
Loading