Skip to content
Closed
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
3 changes: 3 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -141,11 +141,14 @@ docs/*
!docs/COMPOSITE_GROUPS.md
!docs/PLUGIN_DEVELOPMENT.md
!docs/channel-monitor-v2-safe-defaults.md
!docs/codex-credits.md
!docs/legal/
!docs/legal/*.md
!docs/screenshots/
docs/screenshots/*
!docs/screenshots/mobile-account-actions-menu.png
!docs/screenshots/codex-referral-dialog.png
!docs/screenshots/codex-referral-success.png
.serena/
.codex/
frontend/coverage/
Expand Down
11 changes: 10 additions & 1 deletion backend/internal/handler/admin/openai_oauth_handler.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,12 +21,14 @@ type OpenAIOAuthHandler struct {
openaiOAuthService *service.OpenAIOAuthService
adminService service.AdminService
quotaService openAIQuotaService
referralService openAIReferralService
rateLimitService openAIAccountStateRecoverer
}

type openAIQuotaService interface {
QueryUsage(ctx context.Context, accountID int64) (*service.OpenAIQuotaUsage, error)
CacheResetCreditsSnapshot(ctx context.Context, accountID int64, credits *service.OpenAIRateLimitResetCredits) error
CacheCreditsSnapshot(ctx context.Context, accountID int64, usage *service.OpenAIQuotaUsage) error
CachePostResetSnapshot(ctx context.Context, accountID int64, usage *service.OpenAIQuotaUsage) error
ResetCredit(ctx context.Context, accountID int64) (*service.OpenAIQuotaResetResult, error)
}
Expand Down Expand Up @@ -57,7 +59,8 @@ type openAIQuotaResetResponse struct {
// failed display-cache write must never discard a successful upstream read.
type openAIQuotaRefreshResponse struct {
service.OpenAIQuotaUsage
CachePersisted bool `json:"cache_persisted"`
CachePersisted bool `json:"cache_persisted"`
CreditsCachePersisted bool `json:"credits_cache_persisted"`
}

// openAIQuotaResetPostProcessContext detaches the post-reset bookkeeping from the
Expand Down Expand Up @@ -92,6 +95,7 @@ func NewOpenAIOAuthHandler(
// `== nil` capability guards below and panic instead of returning 400.
if quotaService != nil {
h.quotaService = quotaService
h.referralService = quotaService
}
if rateLimitService != nil {
h.rateLimitService = rateLimitService
Expand Down Expand Up @@ -522,6 +526,11 @@ func (h *OpenAIOAuthHandler) RefreshQuota(c *gin.Context) {
service.NotifyOpenAIAutoResetCredit(accountID)

refreshResponse := openAIQuotaRefreshResponse{OpenAIQuotaUsage: *usage}
if err := h.quotaService.CacheCreditsSnapshot(c.Request.Context(), accountID, usage); err != nil {
slog.Warn("openai_quota_credits_cache_persist_failed", "account_id", accountID, "error", err)
} else {
refreshResponse.CreditsCachePersisted = true
}
// A failed snapshot write leaves the previous cache intact — report it as a
// partial success instead of discarding the usage payload we just fetched,
// which would leave the card without a credit count at all.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,11 +17,14 @@ import (
)

type openAIQuotaWorkflowStub struct {
resetResult *service.OpenAIQuotaResetResult
resetErr error
queryResult *service.OpenAIQuotaUsage
queryErr error
cacheErr error
resetResult *service.OpenAIQuotaResetResult
resetErr error
queryResult *service.OpenAIQuotaUsage
queryErr error
cacheErr error
creditsCacheErr error
creditsCacheCalls int
cachedCreditsUsage *service.OpenAIQuotaUsage

resetCalls int
queryCalls int
Expand All @@ -48,6 +51,12 @@ func (s *openAIQuotaWorkflowStub) CacheResetCreditsSnapshot(ctx context.Context,
return s.cacheErr
}

func (s *openAIQuotaWorkflowStub) CacheCreditsSnapshot(_ context.Context, _ int64, usage *service.OpenAIQuotaUsage) error {
s.creditsCacheCalls++
s.cachedCreditsUsage = usage
return s.creditsCacheErr
}

func (s *openAIQuotaWorkflowStub) CachePostResetSnapshot(ctx context.Context, _ int64, _ *service.OpenAIQuotaUsage) error {
s.cacheCalls++
s.cacheCtxErr = ctx.Err()
Expand Down Expand Up @@ -410,6 +419,34 @@ func TestOpenAIRefreshQuota_PersistFailureStillReturnsUsage(t *testing.T) {
require.Equal(t, 1, quota.cacheCalls)
}

func TestOpenAIRefreshQuota_CreditsPersistIndependently(t *testing.T) {
for _, tc := range []struct {
name string
resetErr error
creditsErr error
}{
{name: "both saved"},
{name: "reset details missing", resetErr: errors.New("missing expirations")},
{name: "points cache failed", creditsErr: errors.New("write failed")},
} {
t.Run(tc.name, func(t *testing.T) {
quota := successfulOpenAIQuotaWorkflowStub()
balance := "1250.75"
quota.queryResult.Credits = &service.OpenAICredits{HasCredits: true, Balance: &balance}
quota.cacheErr = tc.resetErr
quota.creditsCacheErr = tc.creditsErr
status, envelope := performOpenAIQuotaRefreshRequest(t, &OpenAIOAuthHandler{quotaService: quota})
require.Equal(t, http.StatusOK, status)
require.Equal(t, tc.resetErr == nil, envelope.Data.CachePersisted)
require.Equal(t, tc.creditsErr == nil, envelope.Data.CreditsCachePersisted)
require.Equal(t, quota.queryResult.Credits, envelope.Data.Credits)
require.Equal(t, quota.queryResult, quota.cachedCreditsUsage)
require.Equal(t, 1, quota.creditsCacheCalls)
require.Zero(t, quota.resetCalls)
})
}
}

// An empty-but-successful upstream read must not be dereferenced blindly.
func TestOpenAIQuotaEmptyUsageIsHandledWithoutPanic(t *testing.T) {
t.Run("refresh reports an internal error", func(t *testing.T) {
Expand Down
96 changes: 96 additions & 0 deletions backend/internal/handler/admin/openai_referral_handler.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,96 @@
package admin

import (
"context"
"net/http"
"strconv"
"time"

"github.com/Wei-Shaw/sub2api/internal/pkg/response"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
)

type openAIReferralService interface {
QueryReferralEligibility(context.Context, int64) (*service.OpenAIReferralEligibility, error)
CacheReferralSnapshot(context.Context, int64, *service.OpenAIReferralEligibility) error
SendReferralInvite(context.Context, int64, service.OpenAIReferralSendRequest) (*service.OpenAIReferralSendResult, error)
}

type openAIReferralRefreshResponse struct {
Eligibility *service.OpenAIReferralEligibility `json:"eligibility"`
CachePersisted bool `json:"cache_persisted"`
}

type openAIReferralSendResponse struct {
service.OpenAIReferralSendResult
openAIReferralRefreshResponse
RefreshFailed bool `json:"refresh_failed"`
}

func (h *OpenAIOAuthHandler) referralAccountID(c *gin.Context) (int64, bool) {
id, err := strconv.ParseInt(c.Param("id"), 10, 64)
if err != nil || id <= 0 {
response.BadRequest(c, "Invalid account ID")
return 0, false
}
if h.referralService == nil {
response.BadRequest(c, "OpenAI referral service is not enabled")
return 0, false
}
return id, true
}

// RefreshReferrals persists a display snapshot, hence POST and admin audit.
func (h *OpenAIOAuthHandler) RefreshReferrals(c *gin.Context) {
id, ok := h.referralAccountID(c)
if !ok {
return
}
eligibility, err := h.referralService.QueryReferralEligibility(c.Request.Context(), id)
if err != nil {
response.ErrorFrom(c, err)
return
}
if eligibility == nil {
response.Error(c, http.StatusBadGateway, "Empty invitation eligibility response")
return
}
cacheErr := h.referralService.CacheReferralSnapshot(c.Request.Context(), id, eligibility)
response.Success(c, openAIReferralRefreshResponse{Eligibility: eligibility, CachePersisted: cacheErr == nil})
}

func (h *OpenAIOAuthHandler) SendReferralInvite(c *gin.Context) {
id, ok := h.referralAccountID(c)
if !ok {
return
}
var input service.OpenAIReferralSendRequest
if err := c.ShouldBindJSON(&input); err != nil {
response.BadRequest(c, "Invalid invitation request")
return
}
result, err := h.referralService.SendReferralInvite(c.Request.Context(), id, input)
if err != nil {
response.ErrorFrom(c, err)
return
}
if result == nil || !result.Sent {
response.Error(c, http.StatusBadGateway, "Invitation outcome is unknown; check Codex before sending again")
return
}
// The email is already sent. Refresh failure must not turn it into a failed
// submission and encourage a duplicate send, even if the browser disconnects.
ctx, cancel := context.WithTimeout(context.WithoutCancel(c.Request.Context()), 8*time.Second)
defer cancel()
eligibility, refreshErr := h.referralService.QueryReferralEligibility(ctx, id)
if refreshErr != nil {
eligibility = nil
}
cacheErr := h.referralService.CacheReferralSnapshot(ctx, id, eligibility)
response.Success(c, openAIReferralSendResponse{
OpenAIReferralSendResult: *result,
openAIReferralRefreshResponse: openAIReferralRefreshResponse{Eligibility: eligibility, CachePersisted: cacheErr == nil},
RefreshFailed: eligibility == nil,
})
}
90 changes: 90 additions & 0 deletions backend/internal/handler/admin/openai_referral_handler_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,90 @@
//go:build unit

package admin

import (
"context"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"strings"
"testing"

"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)

type referralHandlerStub struct {
eligibility *service.OpenAIReferralEligibility
queryErr, cacheErr error
input service.OpenAIReferralSendRequest
cached *service.OpenAIReferralEligibility
sends int
}

func (s *referralHandlerStub) QueryReferralEligibility(context.Context, int64) (*service.OpenAIReferralEligibility, error) {
return s.eligibility, s.queryErr
}
func (s *referralHandlerStub) CacheReferralSnapshot(_ context.Context, _ int64, e *service.OpenAIReferralEligibility) error {
s.cached = e
return s.cacheErr
}
func (s *referralHandlerStub) SendReferralInvite(_ context.Context, _ int64, input service.OpenAIReferralSendRequest) (*service.OpenAIReferralSendResult, error) {
s.sends++
s.input = input
return &service.OpenAIReferralSendResult{Email: input.Email, Sent: true}, nil
}

func referralHandlerRequest(t *testing.T, stub *referralHandlerStub, path, body string) *httptest.ResponseRecorder {
t.Helper()
gin.SetMode(gin.TestMode)
h := &OpenAIOAuthHandler{referralService: stub}
router := gin.New()
router.POST("/:id/referrals/refresh", h.RefreshReferrals)
router.POST("/:id/referrals/invite", h.SendReferralInvite)
rec := httptest.NewRecorder()
router.ServeHTTP(rec, httptest.NewRequest(http.MethodPost, path, strings.NewReader(body)))
return rec
}

func TestOpenAIReferralHandlerRefresh(t *testing.T) {
count := 2
stub := &referralHandlerStub{eligibility: &service.OpenAIReferralEligibility{ShouldShow: true, AvailableInvites: &count}, cacheErr: errors.New("cache failed")}
rec := referralHandlerRequest(t, stub, "/100/referrals/refresh", "")
require.Equal(t, http.StatusOK, rec.Code)
var body struct {
Data openAIReferralRefreshResponse `json:"data"`
}
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &body))
require.Equal(t, 2, *body.Data.Eligibility.AvailableInvites)
require.False(t, body.Data.CachePersisted)
require.Zero(t, stub.sends)
}

func TestOpenAIReferralHandlerSentSurvivesRefreshFailure(t *testing.T) {
stub := &referralHandlerStub{queryErr: errors.New("upstream timed out")}
rec := referralHandlerRequest(t, stub, "/100/referrals/invite", `{"email":"friend@example.com","program_id":"codex_referral_consumer","confirmed":true}`)
require.Equal(t, http.StatusOK, rec.Code)
var body struct {
Data openAIReferralSendResponse `json:"data"`
}
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &body))
require.True(t, body.Data.Sent)
require.Equal(t, "friend@example.com", body.Data.Email)
require.True(t, body.Data.RefreshFailed)
require.True(t, body.Data.CachePersisted)
require.Nil(t, stub.cached, "stale remaining count must be invalidated")
require.True(t, stub.input.Confirmed)
require.Equal(t, 1, stub.sends)
}

func TestOpenAIReferralHandlerRejectsInvalidInput(t *testing.T) {
stub := &referralHandlerStub{}
for _, path := range []string{"/invalid/referrals/invite", "/0/referrals/invite", "/100/referrals/invite"} {
rec := referralHandlerRequest(t, stub, path, "not-json")
require.Equal(t, http.StatusBadRequest, rec.Code)
}
require.Zero(t, stub.sends)
}
2 changes: 2 additions & 0 deletions backend/internal/server/routes/admin.go
Original file line number Diff line number Diff line change
Expand Up @@ -454,6 +454,8 @@ func registerOpenAIOAuthRoutes(admin *gin.RouterGroup, h *handler.Handlers) {
openai.GET("/accounts/:id/quota", h.Admin.OpenAIOAuth.QueryQuota)
openai.POST("/accounts/:id/quota/refresh", h.Admin.OpenAIOAuth.RefreshQuota)
openai.POST("/accounts/:id/reset-quota", h.Admin.OpenAIOAuth.ResetQuota)
openai.POST("/accounts/:id/referrals/refresh", h.Admin.OpenAIOAuth.RefreshReferrals)
openai.POST("/accounts/:id/referrals/invite", h.Admin.OpenAIOAuth.SendReferralInvite)
}
}

Expand Down
Loading