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
5 changes: 3 additions & 2 deletions backend/app/api/routes/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@ def register(body: RegisterRequest, db: Session = Depends(get_db)):

if body.invite_token:
invite = db.query(InviteToken).filter(InviteToken.token == body.invite_token).first()
if not invite or invite.used_by or invite.expires_at < datetime.now():
if not invite or invite.used_by or invite.expires_at < datetime.now(UTC).replace(tzinfo=None):
raise HTTPException(status_code=400, detail="Invalid or expired invite")
if invite.email and invite.email.lower() != body.email.lower():
raise HTTPException(status_code=400, detail="Email does not match invite")
Expand All @@ -46,10 +46,11 @@ def register(body: RegisterRequest, db: Session = Depends(get_db)):
role=role,
)
db.add(user)
db.flush()

if invite:
invite.used_by = user.id
invite.used_at = datetime.now(UTC)
invite.used_at = datetime.now(UTC).replace(tzinfo=None)

db.commit()
db.refresh(user)
Expand Down
4 changes: 2 additions & 2 deletions backend/app/api/routes/invites.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@ def create_invite(
email=body.email,
role=body.role,
created_by=current_user.id,
expires_at=datetime.now(UTC) + timedelta(days=7),
expires_at=datetime.now(UTC).replace(tzinfo=None) + timedelta(days=7),
)
db.add(invite)
db.commit()
Expand Down Expand Up @@ -61,6 +61,6 @@ def delete_invite(
@router.get("/{token}/validate")
def validate_invite(token: str, db: Session = Depends(get_db)):
invite = db.query(InviteToken).filter(InviteToken.token == token).first()
if not invite or invite.used_by or invite.expires_at < datetime.now(UTC):
if not invite or invite.used_by or invite.expires_at < datetime.now(UTC).replace(tzinfo=None):
return {"valid": False}
return {"valid": True, "email": invite.email, "role": invite.role}
114 changes: 114 additions & 0 deletions backend/tests/test_invites.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,114 @@
def test_create_invite(client, auth_header):
resp = client.post(
"/api/invites",
json={"email": "newuser@example.com", "role": "member"},
headers=auth_header,
)
assert resp.status_code == 200
data = resp.json()
assert data["token"]
assert data["email"] == "newuser@example.com"
assert data["role"] == "member"
assert data["used_at"] is None


def test_create_invite_without_email(client, auth_header):
resp = client.post(
"/api/invites",
json={"role": "admin"},
headers=auth_header,
)
assert resp.status_code == 200
data = resp.json()
assert data["email"] is None
assert data["role"] == "admin"


def test_create_invite_requires_admin(client):
client.post(
"/api/auth/register",
json={"email": "first@example.com", "password": "password123", "name": "First"},
)
invite_resp = client.post(
"/api/invites",
json={"email": "new@example.com", "role": "member"},
headers={"Authorization": "Bearer first_token"},
)
assert invite_resp.status_code == 401


def test_validate_invite(client, auth_header):
create_resp = client.post(
"/api/invites",
json={"email": "newuser@example.com", "role": "member"},
headers=auth_header,
)
token = create_resp.json()["token"]

validate_resp = client.get(f"/api/invites/{token}/validate")
assert validate_resp.status_code == 200
data = validate_resp.json()
assert data["valid"] is True
assert data["email"] == "newuser@example.com"
assert data["role"] == "member"


def test_validate_invalid_token(client):
resp = client.get("/api/invites/nonexistent-token/validate")
assert resp.status_code == 200
assert resp.json()["valid"] is False


def test_list_invites(client, auth_header):
client.post(
"/api/invites",
json={"email": "a@example.com"},
headers=auth_header,
)
client.post(
"/api/invites",
json={"email": "b@example.com"},
headers=auth_header,
)
resp = client.get("/api/invites", headers=auth_header)
assert resp.status_code == 200
assert len(resp.json()) == 2


def test_delete_invite(client, auth_header):
create_resp = client.post(
"/api/invites",
json={"email": "delete@example.com"},
headers=auth_header,
)
invite_id = create_resp.json()["id"]

delete_resp = client.delete(f"/api/invites/{invite_id}", headers=auth_header)
assert delete_resp.status_code == 200

list_resp = client.get("/api/invites", headers=auth_header)
assert len(list_resp.json()) == 0


def test_register_with_invite(client, auth_header):
create_resp = client.post(
"/api/invites",
json={"email": "invited@example.com", "role": "member"},
headers=auth_header,
)
token = create_resp.json()["token"]

reg_resp = client.post(
"/api/auth/register",
json={
"email": "invited@example.com",
"password": "password123",
"name": "Invited User",
"invite_token": token,
},
)
assert reg_resp.status_code == 200
assert "access_token" in reg_resp.json()

validate_resp = client.get(f"/api/invites/{token}/validate")
assert validate_resp.json()["valid"] is False
Loading