diff options
| author | Chia <Chia@93.nz> | 2026-08-05 22:01:29 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-05 22:07:50 +1200 |
| commit | eadb2ffe85c43cf6fc741c9823cd28eedb4a844c (patch) | |
| tree | 1aba2536d57360da403aa35c9ced58b615c7064e /internal | |
| parent | cd0dd91ab93653631904f2ea0e574ccde6d60339 (diff) | |
feat: harden prepaid billing and commercial operations
Diffstat (limited to 'internal')
40 files changed, 4080 insertions, 281 deletions
diff --git a/internal/adminapi/api.go b/internal/adminapi/api.go index 4e66d7f..9f460e3 100644 --- a/internal/adminapi/api.go +++ b/internal/adminapi/api.go @@ -144,6 +144,7 @@ func (a *API) Handler() http.Handler { mux.HandleFunc("GET "+apiPrefix+"/models", a.withAuth("platform.read", a.listModels)) mux.HandleFunc("POST "+apiPrefix+"/models", a.withAuth("platform.write", a.createModel)) mux.HandleFunc("POST "+apiPrefix+"/models/{id}/toggle", a.withAuth("platform.write", a.toggleModel)) + mux.HandleFunc("POST "+apiPrefix+"/models/{id}/prices", a.withAuth("platform.write", a.createModelPriceVersion)) mux.HandleFunc("POST "+apiPrefix+"/reload", a.withAuth("platform.write", a.reload)) if a.billing != nil { mux.HandleFunc("GET "+apiPrefix+"/billing/accounts", a.withAuth("billing.read", a.listBillingAccounts)) @@ -152,6 +153,16 @@ func (a *API) Handler() http.Handler { mux.HandleFunc("GET "+apiPrefix+"/billing/orders/{id}", a.withAuth("billing.read", a.getTopUpOrder)) mux.HandleFunc("POST "+apiPrefix+"/billing/adjustments", a.withAuth("billing.adjust", a.adjustBalance)) mux.HandleFunc("POST "+apiPrefix+"/billing/checkout-sessions", a.withAuth("billing.topup", a.createCheckoutSession)) + mux.HandleFunc("POST "+apiPrefix+"/billing/portal-sessions", a.withAuth("billing.topup", a.createPortalSession)) + mux.HandleFunc("POST "+apiPrefix+"/billing/orders/{id}/retry", a.withAuth("billing.topup", a.retryCheckoutSession)) + mux.HandleFunc("POST "+apiPrefix+"/billing/orders/{id}/resolve-missing", a.withAuth("billing.adjust", a.resolveMissingTopUp)) + mux.HandleFunc("POST "+apiPrefix+"/billing/orders/{id}/reverse-missing-credit", a.withAuth("billing.adjust", a.reverseMissingTopUpCredit)) + mux.HandleFunc("POST "+apiPrefix+"/billing/orders/{id}/refund", a.withAuth("billing.adjust", a.createRefund)) + mux.HandleFunc("GET "+apiPrefix+"/billing/refunds", a.withAuth("billing.read", a.listRefunds)) + mux.HandleFunc("GET "+apiPrefix+"/billing/disputes", a.withAuth("billing.read", a.listDisputes)) + mux.HandleFunc("GET "+apiPrefix+"/billing/invoices", a.withAuth("billing.read", a.listInvoices)) + mux.HandleFunc("POST "+apiPrefix+"/billing/reconcile", a.withAuth("billing.adjust", a.reconcileBilling)) + mux.HandleFunc("GET "+apiPrefix+"/billing/export.csv", a.withAuth("billing.read", a.exportBillingCSV)) } mux.HandleFunc("GET "+apiPrefix+"/usage", a.withAuth("usage.read", a.listUsage)) mux.HandleFunc("GET "+apiPrefix+"/usage/summary", a.withAuth("usage.read", a.usageSummary)) @@ -945,11 +956,11 @@ func (a *API) listBillingLedger(w http.ResponseWriter, r *http.Request) { func (a *API) listTopUpOrders(w http.ResponseWriter, r *http.Request) { actor := a.actor(r) - if actor.TenantID == "" { - apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "tenant_required", Message: "A tenant is required"}, requestID(r)) - return + tenantID := actor.TenantID + if tenantID == "" { + tenantID = r.URL.Query().Get("tenant_id") } - result, err := a.billing.ListTopUpOrders(r.Context(), actor.TenantID, 50) + result, err := a.billing.ListTopUpOrders(r.Context(), tenantID, 100) if err != nil { a.billingError(w, r, err) return @@ -959,10 +970,6 @@ func (a *API) listTopUpOrders(w http.ResponseWriter, r *http.Request) { func (a *API) getTopUpOrder(w http.ResponseWriter, r *http.Request) { actor := a.actor(r) - if actor.TenantID == "" { - apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "tenant_required", Message: "A tenant is required"}, requestID(r)) - return - } result, err := a.billing.GetTopUpOrder(r.Context(), actor.TenantID, r.PathValue("id")) if err != nil { a.billingError(w, r, err) @@ -992,6 +999,7 @@ func (a *API) createCheckoutSession(w http.ResponseWriter, r *http.Request) { if tenantID := a.actor(r).TenantID; tenantID != "" { input.TenantID = tenantID } + input.CustomerEmail = a.actor(r).Email result, err := a.billing.CreateCheckout(r.Context(), input) if err != nil { a.billingError(w, r, err) @@ -1000,6 +1008,169 @@ func (a *API) createCheckoutSession(w http.ResponseWriter, r *http.Request) { writeStatusJSON(w, http.StatusCreated, result) } +func (a *API) createPortalSession(w http.ResponseWriter, r *http.Request) { + tenantID := a.actor(r).TenantID + if tenantID == "" { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "tenant_required", Message: "A tenant is required"}, requestID(r)) + return + } + result, err := a.billing.CreatePortalSession(r.Context(), tenantID) + if err != nil { + a.billingError(w, r, err) + return + } + writeStatusJSON(w, http.StatusCreated, result) +} + +func (a *API) retryCheckoutSession(w http.ResponseWriter, r *http.Request) { + actor := a.actor(r) + if actor.TenantID == "" { + a.scopeError(w, r) + return + } + result, err := a.billing.RetryCheckout(r.Context(), actor.TenantID, r.PathValue("id"), actor.Email) + if err != nil { + a.billingError(w, r, err) + return + } + writeStatusJSON(w, http.StatusCreated, result) +} + +func (a *API) resolveMissingTopUp(w http.ResponseWriter, r *http.Request) { + var input billing.ResolveMissingTopUpInput + if !decodeBody(w, r, &input) { + return + } + tenantID := r.URL.Query().Get("tenant_id") + if actorTenant := a.actor(r).TenantID; actorTenant != "" { + tenantID = actorTenant + } + if tenantID == "" { + a.scopeError(w, r) + return + } + actor := a.actor(r) + actorType := "console_user" + if actor.Bootstrap { + actorType = "bootstrap" + } + result, err := a.billing.ResolveMissingTopUp(r.Context(), tenantID, r.PathValue("id"), input, + billing.ResolutionActor{ID: actor.ID, Type: actorType}) + if err != nil { + a.billingError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) reverseMissingTopUpCredit(w http.ResponseWriter, r *http.Request) { + var input billing.ResolveMissingTopUpInput + if !decodeBody(w, r, &input) { + return + } + tenantID := r.URL.Query().Get("tenant_id") + actor := a.actor(r) + if actor.TenantID != "" { + tenantID = actor.TenantID + } + if tenantID == "" { + a.scopeError(w, r) + return + } + actorType := "console_user" + if actor.Bootstrap { + actorType = "bootstrap" + } + result, err := a.billing.ReverseMissingTopUpCredit(r.Context(), tenantID, r.PathValue("id"), input, + billing.ResolutionActor{ID: actor.ID, Type: actorType}) + if err != nil { + a.billingError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) createRefund(w http.ResponseWriter, r *http.Request) { + var input billing.RefundInput + if !decodeBody(w, r, &input) { + return + } + tenantID := r.URL.Query().Get("tenant_id") + if actorTenant := a.actor(r).TenantID; actorTenant != "" { + tenantID = actorTenant + } + if tenantID == "" { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "tenant_required", Message: "tenant_id is required"}, requestID(r)) + return + } + result, err := a.billing.CreateRefund(r.Context(), tenantID, r.PathValue("id"), input) + if err != nil { + a.billingError(w, r, err) + return + } + writeStatusJSON(w, http.StatusAccepted, result) +} + +func (a *API) listRefunds(w http.ResponseWriter, r *http.Request) { + tenantID := a.actor(r).TenantID + if tenantID == "" { + tenantID = r.URL.Query().Get("tenant_id") + } + result, err := a.billing.ListRefunds(r.Context(), tenantID, 200) + if err != nil { + a.billingError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) listDisputes(w http.ResponseWriter, r *http.Request) { + tenantID := a.actor(r).TenantID + if tenantID == "" { + tenantID = r.URL.Query().Get("tenant_id") + } + result, err := a.billing.ListDisputes(r.Context(), tenantID, 200) + if err != nil { + a.billingError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) listInvoices(w http.ResponseWriter, r *http.Request) { + tenantID := a.actor(r).TenantID + if tenantID == "" { + tenantID = r.URL.Query().Get("tenant_id") + } + result, err := a.billing.ListInvoices(r.Context(), tenantID, 200) + if err != nil { + a.billingError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) reconcileBilling(w http.ResponseWriter, r *http.Request) { + result, err := a.billing.Reconcile(r.Context(), 200) + if err != nil { + a.billingError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) exportBillingCSV(w http.ResponseWriter, r *http.Request) { + tenantID := a.actor(r).TenantID + if tenantID == "" { + tenantID = r.URL.Query().Get("tenant_id") + } + w.Header().Set("Content-Type", "text/csv; charset=utf-8") + w.Header().Set("Content-Disposition", `attachment; filename="aigw-financial-ledger.csv"`) + if err := a.billing.WriteFinancialCSV(r.Context(), tenantID, w); err != nil { + a.logger.Error("billing_export_failed", "error", err) + } +} + func (a *API) listTenants(w http.ResponseWriter, r *http.Request) { result, err := a.store.ListTenantsFor(r.Context(), a.actor(r).TenantID) if err != nil { @@ -1185,6 +1356,23 @@ func (a *API) toggleModel(w http.ResponseWriter, r *http.Request) { writeJSON(w, map[string]any{"id": id, "enabled": input.Enabled}) } +func (a *API) createModelPriceVersion(w http.ResponseWriter, r *http.Request) { + var input controlplane.CreatePriceVersionInput + if !decodeBody(w, r, &input) { + return + } + id := r.PathValue("id") + generation, err := a.store.CreateModelPriceVersion(r.Context(), id, input) + if err != nil { + a.mutationError(w, r, err) + return + } + if !a.changed(w, r, generation, "model_price", id) { + return + } + writeStatusJSON(w, http.StatusCreated, map[string]any{"model_id": id, "generation": generation}) +} + func (a *API) reload(w http.ResponseWriter, r *http.Request) { generation, err := a.manager.Reload(r.Context()) if err != nil { @@ -1375,6 +1563,10 @@ func (a *API) billingError(w http.ResponseWriter, r *http.Request, err error) { status = http.StatusNotFound typeName = "topup_order_not_found" message = "Top-up order was not found" + case errors.Is(err, billing.ErrCannotResolveTopUp): + status = http.StatusConflict + typeName = "topup_order_not_resolvable" + message = err.Error() default: a.logger.Error("admin_billing_error", "error", err) } diff --git a/internal/adminui/assets/app.js b/internal/adminui/assets/app.js index 1a965fc..8d29e23 100644 --- a/internal/adminui/assets/app.js +++ b/internal/adminui/assets/app.js @@ -1,7 +1,7 @@ const state = { token: '', csrf: '', actor: {}, permissions: new Set(), overview: {}, tenants: [], projects: [], keys: [], providers: [], models: [], billingAccounts: [], ledger: [], - usage: [], usageSummary: [], limits: [], users: [], audit: [], orders: [], sessions: [], + usage: [], usageSummary: [], limits: [], users: [], audit: [], orders: [], refunds: [], disputes: [], invoices: [], sessions: [], mfa: {totp_enabled:false,passkeys:[]}, pendingMFA: null, authConfig: {} }; const $ = (selector) => document.querySelector(selector); @@ -102,9 +102,12 @@ async function loadAll(knownSession = null) { state.overview.billing_enabled ? permitted('billing.read','/billing/ledger') : [], state.actor.id ? api('/auth/mfa') : {totp_enabled:false,passkeys:[]}, state.actor.id ? api('/auth/sessions') : [], - state.actor.tenant_id && state.overview.billing_enabled && can('billing.read') ? api('/billing/orders') : [] + state.overview.billing_enabled && can('billing.read') ? api('/billing/orders') : [], + state.overview.billing_enabled ? permitted('billing.read','/billing/refunds') : [], + state.overview.billing_enabled ? permitted('billing.read','/billing/disputes') : [], + state.overview.billing_enabled ? permitted('billing.read','/billing/invoices') : [] ]); - [state.tenants,state.projects,state.keys,state.providers,state.models,state.usage,state.usageSummary,state.limits,state.users,state.audit,state.billingAccounts,state.ledger,state.mfa,state.sessions,state.orders] = results; + [state.tenants,state.projects,state.keys,state.providers,state.models,state.usage,state.usageSummary,state.limits,state.users,state.audit,state.billingAccounts,state.ledger,state.mfa,state.sessions,state.orders,state.refunds,state.disputes,state.invoices] = results; renderAll(); setConnected(true); return true; } catch (error) { setConnected(false); if (error.status !== 401) toast(error.message, true); return false; } } @@ -137,11 +140,14 @@ function renderProjects() { $('#projects-body').innerHTML = state.projects.map(i function renderKeyProjects() { const tenant = $('#key-tenant').value; const projects = state.projects.filter(item => !tenant || item.tenant_id === tenant); $('#key-project').innerHTML = selectOptions(projects,'id','name'); } function renderKeys() { $('#keys-body').innerHTML = state.keys.map(item => `<tr><td><strong>${esc(item.name)}</strong></td><td><code>${esc(item.key_prefix)}</code></td><td><code>${shortID(item.project_id)}</code></td><td>${(item.scopes||[]).map(scope=>`<span class="tag">${esc(scope)}</span>`).join('')}</td><td><span class="badge ${item.status}">${esc(item.status)}</span></td><td>${item.status==='active'&&can('keys.write')?`<button class="text-button danger" data-revoke-key="${esc(item.id)}">Revoke</button>`:''}</td></tr>`).join('') || emptyRow(6); } function renderProviders() { $('#providers-body').innerHTML = state.providers.map(item => `<tr><td><strong>${esc(item.name)}</strong></td><td><span class="tag">${esc(item.protocol)}</span></td><td class="truncate">${esc(item.base_url)}</td><td>${integer(item.route_count)}</td><td><span class="badge ${item.enabled?'active':'suspended'}">${item.enabled?'enabled':'disabled'}</span></td><td>${can('platform.write')?`<button class="text-button" data-toggle-provider="${esc(item.id)}" data-enabled="${!item.enabled}">${item.enabled?'Disable':'Enable'}</button>`:''}</td></tr>`).join('') || emptyRow(6); } -function renderModels() { $('#models-body').innerHTML = state.models.map(item => `<tr><td><strong>${esc(item.public_id)}</strong></td><td>${esc(item.owned_by||'—')}<small class="price-line">in ${money(item.input_price_micros_per_million)}/1M · out ${money(item.output_price_micros_per_million)}/1M</small></td><td><div class="route-list">${(item.routes||[]).map(route=>`<span>${esc(route.provider_name||route.provider_id).slice(0,24)} → ${esc(route.upstream_model)} <em>p${route.priority} / w${route.weight}</em></span>`).join('')}</div></td><td><span class="badge ${item.enabled?'active':'suspended'}">${item.enabled?'enabled':'disabled'}</span></td><td>${can('platform.write')?`<button class="text-button" data-toggle-model="${esc(item.id)}" data-enabled="${!item.enabled}">${item.enabled?'Disable':'Enable'}</button>`:''}</td></tr>`).join('') || emptyRow(5); } +function renderModels() { $('#models-body').innerHTML = state.models.map(item => `<tr><td><strong>${esc(item.public_id)}</strong><small class="price-line">${esc(item.display_name||'')} · ${integer(item.context_window)} ctx · ${esc((item.input_modalities||[]).join('+'))} → ${esc((item.output_modalities||[]).join('+'))}</small></td><td>${esc(item.owned_by||'—')}<small class="price-line">v${item.price_version||1} ${esc(item.price_currency||'usd')} · in ${money(item.input_price_micros_per_million,item.price_currency)}/1M · out ${money(item.output_price_micros_per_million,item.price_currency)}/1M</small></td><td><div class="route-list">${(item.routes||[]).map(route=>`<span>${esc(route.provider_name||route.provider_id).slice(0,24)} → ${esc(route.upstream_model)} <em>p${route.priority} / w${route.weight}</em></span>`).join('')}</div></td><td><span class="badge ${item.enabled&&item.lifecycle!=='retired'?'active':'suspended'}">${esc(item.lifecycle||'active')}</span></td><td>${can('platform.write')?`<button class="text-button" data-toggle-model="${esc(item.id)}" data-enabled="${!item.enabled}">${item.enabled?'Disable':'Enable'}</button>`:''}</td></tr>`).join('') || emptyRow(5); } function renderBilling() { $('#billing-currency').textContent=(state.overview.billing_currency||'').toUpperCase(); $('#billing-accounts-body').innerHTML=state.billingAccounts.map(item=>`<tr><td><strong>${esc(item.tenant_name)}</strong><br><code>${shortID(item.tenant_id)}</code></td><td>${money(item.balance_micros,item.currency)}</td><td>${money(item.reserved_micros,item.currency)}</td><td><strong>${money(item.available_micros,item.currency)}</strong></td><td>${date(item.updated_at)}</td></tr>`).join('')||emptyRow(5); $('#billing-ledger-body').innerHTML=state.ledger.map(item=>`<tr><td>${date(item.created_at)}</td><td><code>${shortID(item.tenant_id)}</code></td><td><span class="tag">${esc(item.kind)}</span></td><td class="${item.amount_micros>=0?'money-positive':'money-negative'}">${money(item.amount_micros,item.currency)}</td><td>${money(item.balance_after_micros,item.currency)}</td><td title="${esc(item.description)}"><code>${shortID(item.source_id)}</code></td></tr>`).join('')||emptyRow(6); + const orders=state.orders||[];$('#billing-orders-body').innerHTML=orders.map(item=>`<tr><td>${date(item.created_at)}</td><td>${money(item.amount_micros,item.currency)}</td><td><span class="badge ${item.status==='paid'?'active':'suspended'}">${esc(item.status)}</span></td><td title="${esc(item.reconciliation_error||'')}"><span class="badge ${['ok','repaired','resolved'].includes(item.reconciliation_status)?'active':'suspended'}">${esc(item.reconciliation_status||'unknown')}</span></td><td>${item.invoice_url?`<a href="${esc(item.invoice_url)}" target="_blank" rel="noopener">Invoice</a>`:''} ${item.invoice_pdf_url?`<a href="${esc(item.invoice_pdf_url)}" target="_blank" rel="noopener">PDF</a>`:''} ${item.receipt_url?`<a href="${esc(item.receipt_url)}" target="_blank" rel="noopener">Receipt</a>`:''}</td><td>${['failed','expired'].includes(item.status)?`<button class="text-button" data-retry-order="${esc(item.id)}">Retry</button>`:''}${can('billing.adjust')&&item.status==='pending'&&item.reconciliation_status==='missing'?` <button class="text-button danger" data-resolve-order="${esc(item.id)}" data-order-tenant="${esc(item.tenant_id)}">Resolve</button>`:''}${can('billing.adjust')&&item.status==='paid'&&item.reconciliation_status==='missing'&&!item.stripe_payment_intent_id?` <button class="text-button danger" data-reverse-order="${esc(item.id)}" data-order-tenant="${esc(item.tenant_id)}">Reverse</button>`:''}${can('billing.adjust')&&['paid','partially_refunded'].includes(item.status)?` <button class="text-button danger" data-refund-order="${esc(item.id)}" data-order-amount="${item.amount_minor}">Refund</button>`:''}</td></tr>`).join('')||emptyRow(6); + $('#refunds-body').innerHTML=(state.refunds||[]).map(item=>`<tr><td>${date(item.created_at)}</td><td><code>${shortID(item.topup_order_id)}</code></td><td>${money(item.amount_micros,item.currency)}</td><td><span class="badge ${item.status==='succeeded'?'active':'suspended'}">${esc(item.status)}</span></td><td>${esc(item.last_error||'—')}</td></tr>`).join('')||emptyRow(5); + $('#disputes-body').innerHTML=(state.disputes||[]).map(item=>`<tr><td>${date(item.updated_at)}</td><td>${money(item.amount_micros,item.currency)}</td><td><span class="badge ${item.status==='won'?'active':'suspended'}">${esc(item.status)}</span></td><td>${esc(item.reason)}</td><td>${date(item.due_by)}</td></tr>`).join('')||emptyRow(5); } function renderSecurity() { const totp=Boolean(state.mfa?.totp_enabled); @@ -182,6 +188,13 @@ document.addEventListener('click',async(event)=>{ const session=event.target.closest('[data-revoke-session]');if(session&&confirm('Sign out this device?')){try{const result=await api(`/auth/sessions/${session.dataset.revokeSession}/revoke`,{method:'POST',body:'{}'});if(result.current){state.csrf='';state.actor={};setConnected(false);}else{await loadAll();toast('Device signed out');}}catch(error){toast(error.message,true);}} const passkey=event.target.closest('[data-delete-passkey]');if(passkey&&confirm('Delete this passkey?')){const password=$('#passkey-form [name=current_password]').value;if(!password){toast('Enter your current password in the Passkey panel',true);}else{try{await api(`/auth/passkeys/${passkey.dataset.deletePasskey}/delete`,{method:'POST',body:JSON.stringify({current_password:password})});await loadAll();toast('Passkey deleted');}catch(error){toast(error.message,true);}}} const save=event.target.closest('.save-limit');if(save){const row=save.closest('[data-limit-project]');try{await api(`/limits/${row.dataset.limitProject}`,{method:'POST',body:JSON.stringify({requests_per_minute:Number(row.querySelector('.limit-rpm').value),tokens_per_minute:Number(row.querySelector('.limit-tpm').value),concurrent_requests:Number(row.querySelector('.limit-concurrency').value),monthly_spend_micros:decimalToScaled(row.querySelector('.limit-spend').value,6)})});await loadAll();toast('Project limits updated');}catch(error){toast(error.message,true);}} + if(event.target.id==='billing-portal'){try{const result=await api('/billing/portal-sessions',{method:'POST',body:'{}'});window.location.assign(result.url);}catch(error){toast(error.message,true);}} + if(event.target.id==='billing-export'){try{const headers={};if(state.token)headers.Authorization=`Bearer ${state.token}`;const response=await fetch('./api/billing/export.csv',{credentials:'same-origin',headers});if(!response.ok){const payload=await response.json().catch(()=>({}));throw new Error(payload?.error?.message||`Export failed (${response.status})`);}const blob=await response.blob();const url=URL.createObjectURL(blob);const link=document.createElement('a');link.href=url;link.download=`aigw-ledger-${new Date().toISOString().slice(0,10)}.csv`;document.body.appendChild(link);link.click();link.remove();URL.revokeObjectURL(url);toast('Financial export downloaded');}catch(error){toast(error.message,true);}} + if(event.target.id==='billing-reconcile'){try{const result=await api('/billing/reconcile',{method:'POST',body:'{}'});toast(`Reconciliation ${result.status}: ${result.mismatch_count} mismatches`,result.mismatch_count>0);await loadAll();}catch(error){toast(error.message,true);}} + const retryOrder=event.target.closest('[data-retry-order]');if(retryOrder){try{const result=await api(`/billing/orders/${retryOrder.dataset.retryOrder}/retry`,{method:'POST',body:'{}'});window.location.assign(result.url);}catch(error){toast(error.message,true);}} + const resolveOrder=event.target.closest('[data-resolve-order]');if(resolveOrder){const reason=prompt('Resolution reason for this missing Stripe Session');if(reason){try{await api(`/billing/orders/${resolveOrder.dataset.resolveOrder}/resolve-missing?tenant_id=${encodeURIComponent(resolveOrder.dataset.orderTenant)}`,{method:'POST',body:JSON.stringify({reason})});await loadAll();toast('Missing top-up resolved');}catch(error){toast(error.message,true);}}} + const reverseOrder=event.target.closest('[data-reverse-order]');if(reverseOrder){const reason=prompt('Reason for reversing this unverified local credit');if(reason&&confirm('Post an equal negative ledger entry for this credit?')){try{await api(`/billing/orders/${reverseOrder.dataset.reverseOrder}/reverse-missing-credit?tenant_id=${encodeURIComponent(reverseOrder.dataset.orderTenant)}`,{method:'POST',body:JSON.stringify({reason})});await loadAll();toast('Unverified credit reversed');}catch(error){toast(error.message,true);}}} + const refundOrder=event.target.closest('[data-refund-order]');if(refundOrder){const amount=prompt('Refund amount in account currency');if(amount){try{const digits=currencyDigits(state.overview.billing_currency||'usd');await api(`/billing/orders/${refundOrder.dataset.refundOrder}/refund?tenant_id=${encodeURIComponent(state.orders.find(item=>item.id===refundOrder.dataset.refundOrder)?.tenant_id||'')}`,{method:'POST',body:JSON.stringify({amount_minor:decimalToScaled(amount,digits),reason:'requested_by_customer'})});await loadAll();toast('Refund queued');}catch(error){toast(error.message,true);}}} }); $('#key-tenant').addEventListener('change',renderKeyProjects); @@ -200,7 +213,7 @@ $('#tenant-form').addEventListener('submit',async(event)=>{event.preventDefault( $('#project-form').addEventListener('submit',async(event)=>{event.preventDefault();try{await api('/projects',{method:'POST',body:JSON.stringify(formJSON(event.target))});event.target.reset();await loadAll();toast('Project created');}catch(error){toast(error.message,true);}}); $('#key-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);data.scopes=data.scopes.split(',').map(value=>value.trim()).filter(Boolean);const result=await api('/keys',{method:'POST',body:JSON.stringify(data)});event.target.reset();showSecret('API key created',result.key);await loadAll();}catch(error){toast(error.message,true);}}); $('#provider-form').addEventListener('submit',async(event)=>{event.preventDefault();try{await api('/providers',{method:'POST',body:JSON.stringify(formJSON(event.target))});event.target.reset();await loadAll();toast('Provider added');}catch(error){toast(error.message,true);}}); -$('#model-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);data.input_price_micros_per_million=decimalToScaled(data.input_price,6);data.output_price_micros_per_million=decimalToScaled(data.output_price,6);data.cache_read_price_micros_per_million=decimalToScaled(data.cache_read_price,6);data.cache_write_price_micros_per_million=decimalToScaled(data.cache_write_price,6);delete data.input_price;delete data.output_price;delete data.cache_read_price;delete data.cache_write_price;data.routes=$$('.route-row').map(row=>({provider_id:row.querySelector('.route-provider').value,upstream_model:row.querySelector('.route-upstream').value,priority:Number(row.querySelector('.route-priority').value),weight:Number(row.querySelector('.route-weight').value)}));await api('/models',{method:'POST',body:JSON.stringify(data)});event.target.reset();$('#route-editor').innerHTML='';renderRouteEditor();await loadAll();toast('Model created');}catch(error){toast(error.message,true);}}); +$('#model-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);data.input_price_micros_per_million=decimalToScaled(data.input_price,6);data.output_price_micros_per_million=decimalToScaled(data.output_price,6);data.cache_read_price_micros_per_million=decimalToScaled(data.cache_read_price,6);data.cache_write_price_micros_per_million=decimalToScaled(data.cache_write_price,6);delete data.input_price;delete data.output_price;delete data.cache_read_price;delete data.cache_write_price;for(const field of ['capabilities','input_modalities','output_modalities','regions','aliases','allowed_tenant_ids','allowed_key_ids'])data[field]=String(data[field]||'').split(',').map(value=>value.trim()).filter(Boolean);data.context_window=Number(data.context_window||0);data.max_output_tokens=Number(data.max_output_tokens||0);data.price_currency=state.overview.billing_currency||'usd';data.routes=$$('.route-row').map(row=>({provider_id:row.querySelector('.route-provider').value,upstream_model:row.querySelector('.route-upstream').value,priority:Number(row.querySelector('.route-priority').value),weight:Number(row.querySelector('.route-weight').value)}));await api('/models',{method:'POST',body:JSON.stringify(data)});event.target.reset();$('#route-editor').innerHTML='';renderRouteEditor();await loadAll();toast('Model created');}catch(error){toast(error.message,true);}}); $('#topup-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);const digits=currencyDigits(state.overview.billing_currency||'usd');const result=await api('/billing/checkout-sessions',{method:'POST',body:JSON.stringify({tenant_id:data.tenant_id,amount_minor:decimalToScaled(data.amount,digits)})});window.location.assign(result.url);}catch(error){toast(error.message,true);}}); $('#adjustment-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);await api('/billing/adjustments',{method:'POST',body:JSON.stringify({tenant_id:data.tenant_id,amount_micros:decimalToScaled(data.amount,6),description:data.description})});event.target.reset();await loadAll();toast('Balance adjusted');}catch(error){toast(error.message,true);}}); $('#user-form').addEventListener('submit',async(event)=>{event.preventDefault();try{await api('/users',{method:'POST',body:JSON.stringify(formJSON(event.target))});event.target.reset();await loadAll();toast('Invitation sent');}catch(error){toast(error.message,true);}}); diff --git a/internal/adminui/assets/index.html b/internal/adminui/assets/index.html index 95ab153..04b4e87 100644 --- a/internal/adminui/assets/index.html +++ b/internal/adminui/assets/index.html @@ -46,14 +46,14 @@ <form class="auth-pane" id="reset-complete-pane"> <div><span class="eyebrow">ACCOUNT RECOVERY</span><h1>Choose a new password</h1></div> <input name="token" id="reset-token" type="hidden"> - <input class="visually-hidden" autocomplete="username" aria-hidden="true" tabindex="-1"> + <input class="visually-hidden" name="username" autocomplete="username" aria-hidden="true" tabindex="-1"> <label>New password<input name="new_password" type="password" minlength="12" maxlength="128" required autocomplete="new-password"></label> <button class="button primary" type="submit">Update password</button> </form> <form class="auth-pane" id="invite-pane"> <div><span class="eyebrow">TEAM ACCESS</span><h1>Accept invitation</h1></div> <input name="token" id="invite-token" type="hidden"> - <input class="visually-hidden" autocomplete="username" aria-hidden="true" tabindex="-1"> + <input class="visually-hidden" name="username" autocomplete="username" aria-hidden="true" tabindex="-1"> <label>Name<input name="display_name" required autocomplete="name"></label> <label>Password<input name="password" type="password" minlength="12" maxlength="128" required autocomplete="new-password"></label> <button class="button primary" type="submit">Join workspace</button> @@ -124,18 +124,18 @@ <section id="providers" class="section"> <div class="section-heading"><div><span class="eyebrow">UPSTREAMS</span><h1>Providers</h1></div></div> - <form class="panel form-grid" id="provider-form" data-permission="platform.write"><input class="visually-hidden" autocomplete="username" value="aigw-provider" aria-hidden="true" tabindex="-1"><label>Name<input name="name" required placeholder="openai-primary"></label><label>Protocol<select name="protocol"><option value="openai">OpenAI</option><option value="anthropic">Anthropic</option></select></label><label>Base URL<input name="base_url" type="url" required placeholder="https://api.example.com/v1"></label><label>API key<input name="api_key" type="password" required autocomplete="new-password" placeholder="Stored encrypted"></label><button class="button primary" type="submit">Add provider</button></form> + <form class="panel form-grid" id="provider-form" data-permission="platform.write"><input class="visually-hidden" name="username" autocomplete="username" value="aigw-provider" aria-hidden="true" tabindex="-1"><label>Name<input name="name" autocomplete="off" required placeholder="openai-primary"></label><label>Protocol<select name="protocol"><option value="openai">OpenAI</option><option value="anthropic">Anthropic</option></select></label><label>Base URL<input name="base_url" type="url" autocomplete="off" required placeholder="https://api.example.com/v1"></label><label>API key<input name="api_key" type="password" required autocomplete="new-password" placeholder="Stored encrypted"></label><button class="button primary" type="submit">Add provider</button></form> <div class="panel table-wrap"><table><thead><tr><th>Name</th><th>Protocol</th><th>Base URL</th><th>Routes</th><th>Status</th><th></th></tr></thead><tbody id="providers-body"></tbody></table></div> </section> <section id="models" class="section"> <div class="section-heading"><div><span class="eyebrow">ROUTING</span><h1>Models & routes</h1></div></div> - <form class="panel form-grid" id="model-form" data-permission="platform.write"><label>Public model ID<input name="public_id" required placeholder="openai/gpt-4.1-mini"></label><label>Owned by<input name="owned_by" placeholder="openai"></label><label>Input price / 1M<input name="input_price" inputmode="decimal" value="0" required></label><label>Output price / 1M<input name="output_price" inputmode="decimal" value="0" required></label><label>Cache read / 1M<input name="cache_read_price" inputmode="decimal" value="0" required></label><label>Cache write / 1M<input name="cache_write_price" inputmode="decimal" value="0" required></label><div class="route-editor" id="route-editor"></div><button class="button subtle" type="button" id="add-route">Add route</button><button class="button primary" type="submit">Create model</button></form> + <form class="panel form-grid" id="model-form" data-permission="platform.write"><label>Public model ID<input name="public_id" required placeholder="openai/gpt-4.1-mini"></label><label>Display name<input name="display_name" placeholder="GPT 4.1 mini"></label><label>Owned by<input name="owned_by" placeholder="openai"></label><label>Lifecycle<select name="lifecycle"><option value="preview">Preview</option><option value="active" selected>Active</option><option value="deprecated">Deprecated</option><option value="retired">Retired</option></select></label><label>Description<input name="description" placeholder="Fast general-purpose text model"></label><label>Context window<input name="context_window" type="number" min="0" value="0"></label><label>Max output tokens<input name="max_output_tokens" type="number" min="0" value="0"></label><label>Capabilities<input name="capabilities" value="chat,streaming" placeholder="chat,streaming,tools,json"></label><label>Input modalities<input name="input_modalities" value="text" placeholder="text,image"></label><label>Output modalities<input name="output_modalities" value="text" placeholder="text,image"></label><label>Regions<input name="regions" placeholder="us-east,nz"></label><label>Deprecated aliases<input name="aliases" placeholder="old/model-id"></label><label>Tenant allowlist IDs<input name="allowed_tenant_ids" placeholder="UUIDs, blank is public"></label><label>API key allowlist IDs<input name="allowed_key_ids" placeholder="UUIDs, blank is unrestricted"></label><label>Input price / 1M<input name="input_price" inputmode="decimal" value="0" required></label><label>Output price / 1M<input name="output_price" inputmode="decimal" value="0" required></label><label>Cache read / 1M<input name="cache_read_price" inputmode="decimal" value="0" required></label><label>Cache write / 1M<input name="cache_write_price" inputmode="decimal" value="0" required></label><div class="route-editor" id="route-editor"></div><button class="button subtle" type="button" id="add-route">Add route</button><button class="button primary" type="submit">Create model</button></form> <div class="panel table-wrap"><table><thead><tr><th>Public ID</th><th>Owner</th><th>Routes</th><th>Status</th><th></th></tr></thead><tbody id="models-body"></tbody></table></div> </section> <section id="billing" class="section"> - <div class="section-heading"><div><span class="eyebrow">REVENUE</span><h1>Balances & ledger</h1></div><span class="currency-label" id="billing-currency"></span></div> + <div class="section-heading"><div><span class="eyebrow">REVENUE</span><h1>Balances & ledger</h1></div><div class="form-actions"><button class="button secondary" id="billing-portal" data-permission="billing.topup">Customer portal</button><button class="button subtle" id="billing-export" data-permission="billing.read">Export CSV</button><button class="button subtle" id="billing-reconcile" data-permission="billing.adjust">Reconcile Stripe</button><span class="currency-label" id="billing-currency"></span></div></div> <div class="billing-actions"> <form class="panel form-grid compact-form" id="topup-form" data-permission="billing.topup"><label>Tenant<select name="tenant_id" id="topup-tenant" required></select></label><label>Amount<input name="amount" inputmode="decimal" min="0" required placeholder="25.00"></label><button class="button primary" type="submit">Open Stripe Checkout</button></form> <form class="panel form-grid compact-form" id="adjustment-form" data-permission="billing.adjust"><label>Tenant<select name="tenant_id" id="adjustment-tenant" required></select></label><label>Signed amount<input name="amount" inputmode="decimal" required placeholder="10.00 or -5.00"></label><label>Reference<input name="description" maxlength="240" placeholder="Support credit"></label><button class="button secondary" type="submit">Post adjustment</button></form> @@ -143,6 +143,10 @@ <div class="panel table-wrap"><table><thead><tr><th>Tenant</th><th>Balance</th><th>Reserved</th><th>Available</th><th>Updated</th></tr></thead><tbody id="billing-accounts-body"></tbody></table></div> <div class="section-heading ledger-heading"><div><span class="eyebrow">AUDIT</span><h2>Recent ledger entries</h2></div></div> <div class="panel table-wrap"><table><thead><tr><th>Time</th><th>Tenant</th><th>Kind</th><th>Amount</th><th>Balance after</th><th>Reference</th></tr></thead><tbody id="billing-ledger-body"></tbody></table></div> + <div class="section-heading ledger-heading"><div><span class="eyebrow">PAYMENTS</span><h2>Top-ups, invoices & refunds</h2></div></div> + <div class="panel table-wrap"><table><thead><tr><th>Created</th><th>Amount</th><th>Status</th><th>Reconciliation</th><th>Documents</th><th>Action</th></tr></thead><tbody id="billing-orders-body"></tbody></table></div> + <div class="panel table-wrap"><table><thead><tr><th>Created</th><th>Order</th><th>Amount</th><th>Status</th><th>Failure</th></tr></thead><tbody id="refunds-body"></tbody></table></div> + <div class="panel table-wrap"><table><thead><tr><th>Updated</th><th>Amount</th><th>Status</th><th>Reason</th><th>Due</th></tr></thead><tbody id="disputes-body"></tbody></table></div> </section> <section id="limits" class="section"> @@ -160,9 +164,9 @@ <div class="section-heading"><div><span class="eyebrow">SECURITY</span><h1>Account</h1></div></div> <form class="panel form-grid compact-form" id="password-form"><input class="visually-hidden" name="username" autocomplete="username" aria-hidden="true" tabindex="-1"><label>Current password<input name="current_password" type="password" required autocomplete="current-password"></label><label>New password<input name="new_password" type="password" required minlength="12" maxlength="128" autocomplete="new-password"></label><button class="button primary" type="submit">Change password</button></form> <div class="account-grid"> - <form class="panel form-grid compact-form" id="totp-begin-form"><h2>Authenticator app</h2><input class="visually-hidden" autocomplete="username" aria-hidden="true" tabindex="-1"><label>Current password<input name="current_password" type="password" required autocomplete="current-password"></label><label class="hidden" id="totp-disable-code">Authenticator code<input name="code" inputmode="numeric" autocomplete="one-time-code"></label><button class="button secondary" id="totp-action" type="submit">Set up TOTP</button></form> + <form class="panel form-grid compact-form" id="totp-begin-form"><h2>Authenticator app</h2><input class="visually-hidden" name="username" autocomplete="username" aria-hidden="true" tabindex="-1"><label>Current password<input name="current_password" type="password" required autocomplete="current-password"></label><label class="hidden" id="totp-disable-code">Authenticator code<input name="code" inputmode="numeric" autocomplete="one-time-code"></label><button class="button secondary" id="totp-action" type="submit">Set up TOTP</button></form> <form class="panel form-grid compact-form hidden" id="totp-confirm-form"><h2>Confirm authenticator</h2><img id="totp-qr" alt="TOTP QR code"><code id="totp-secret"></code><label>Code<input name="code" inputmode="numeric" autocomplete="one-time-code" required></label><button class="button primary" type="submit">Enable TOTP</button></form> - <form class="panel form-grid compact-form" id="passkey-form"><h2>Passkey</h2><input class="visually-hidden" autocomplete="username" aria-hidden="true" tabindex="-1"><label>Current password<input name="current_password" type="password" required autocomplete="current-password"></label><label>Passkey name<input name="name" value="This device" required></label><button class="button secondary" type="submit">Add passkey</button></form> + <form class="panel form-grid compact-form" id="passkey-form"><h2>Passkey</h2><input class="visually-hidden" name="username" autocomplete="username" aria-hidden="true" tabindex="-1"><label>Current password<input name="current_password" type="password" required autocomplete="current-password"></label><label>Passkey name<input name="name" autocomplete="off" value="This device" required></label><button class="button secondary" type="submit">Add passkey</button></form> </div> <div class="panel table-wrap"><div class="section-heading"><h2>Devices</h2><button class="button subtle" id="revoke-other-sessions" type="button">Sign out other devices</button></div><table><thead><tr><th>Device</th><th>IP</th><th>Method</th><th>Last seen</th><th></th></tr></thead><tbody id="sessions-body"></tbody></table></div> <div class="panel table-wrap"><div class="section-heading"><h2>Passkeys</h2></div><table><thead><tr><th>Name</th><th>Created</th><th>Last used</th><th></th></tr></thead><tbody id="passkeys-body"></tbody></table></div> diff --git a/internal/billing/ledger.go b/internal/billing/ledger.go index c408cfe..2eb3d87 100644 --- a/internal/billing/ledger.go +++ b/internal/billing/ledger.go @@ -77,9 +77,20 @@ func (s *Service) ListTopUpOrders(ctx context.Context, tenantID string, limit in if limit < 1 || limit > 200 { limit = 50 } - rows, err := s.db.Query(ctx, `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status, - COALESCE(stripe_session_id,''),COALESCE(checkout_url,''),created_at,paid_at FROM topup_orders - WHERE tenant_id=$1 ORDER BY created_at DESC LIMIT $2`, tenantID, limit) + query := `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status, + COALESCE(stripe_session_id,''),COALESCE(checkout_url,''),created_at,paid_at, + COALESCE(stripe_customer_id,''),COALESCE(stripe_payment_intent_id,''),COALESCE(stripe_charge_id,''), + COALESCE(stripe_invoice_id,''),COALESCE(invoice_url,''),COALESCE(invoice_pdf_url,''),COALESCE(receipt_url,''), + refunded_micros,disputed_micros,reconciliation_status,reconciled_at,reconciliation_error FROM topup_orders` + args := []any{} + if strings.TrimSpace(tenantID) != "" { + query += ` WHERE tenant_id=$1 ORDER BY created_at DESC LIMIT $2` + args = []any{tenantID, limit} + } else { + query += ` ORDER BY created_at DESC LIMIT $1` + args = []any{limit} + } + rows, err := s.db.Query(ctx, query, args...) if err != nil { return nil, fmt.Errorf("query top-up orders: %w", err) } @@ -88,7 +99,10 @@ func (s *Service) ListTopUpOrders(ctx context.Context, tenantID string, limit in for rows.Next() { var item TopUpOrder if err := rows.Scan(&item.ID, &item.TenantID, &item.AmountMinor, &item.AmountMicros, &item.Currency, &item.Status, - &item.StripeSessionID, &item.CheckoutURL, &item.CreatedAt, &item.PaidAt); err != nil { + &item.StripeSessionID, &item.CheckoutURL, &item.CreatedAt, &item.PaidAt, &item.StripeCustomerID, + &item.StripePaymentIntentID, &item.StripeChargeID, &item.StripeInvoiceID, &item.InvoiceURL, + &item.InvoicePDFURL, &item.ReceiptURL, &item.RefundedMicros, &item.DisputedMicros, + &item.ReconciliationStatus, &item.ReconciledAt, &item.ReconciliationError); err != nil { return nil, err } result = append(result, item) @@ -98,10 +112,21 @@ func (s *Service) ListTopUpOrders(ctx context.Context, tenantID string, limit in func (s *Service) GetTopUpOrder(ctx context.Context, tenantID, orderID string) (TopUpOrder, error) { var result TopUpOrder - err := s.db.QueryRow(ctx, `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status, - COALESCE(stripe_session_id,''),COALESCE(checkout_url,''),created_at,paid_at FROM topup_orders - WHERE id=$1 AND tenant_id=$2`, orderID, tenantID).Scan(&result.ID, &result.TenantID, &result.AmountMinor, &result.AmountMicros, - &result.Currency, &result.Status, &result.StripeSessionID, &result.CheckoutURL, &result.CreatedAt, &result.PaidAt) + query := `SELECT id::text,tenant_id::text,amount_minor,amount_micros,currency,status, + COALESCE(stripe_session_id,''),COALESCE(checkout_url,''),created_at,paid_at, + COALESCE(stripe_customer_id,''),COALESCE(stripe_payment_intent_id,''),COALESCE(stripe_charge_id,''), + COALESCE(stripe_invoice_id,''),COALESCE(invoice_url,''),COALESCE(invoice_pdf_url,''),COALESCE(receipt_url,''), + refunded_micros,disputed_micros,reconciliation_status,reconciled_at,reconciliation_error FROM topup_orders WHERE id=$1` + args := []any{orderID} + if strings.TrimSpace(tenantID) != "" { + query += ` AND tenant_id=$2` + args = append(args, tenantID) + } + err := s.db.QueryRow(ctx, query, args...).Scan(&result.ID, &result.TenantID, &result.AmountMinor, &result.AmountMicros, + &result.Currency, &result.Status, &result.StripeSessionID, &result.CheckoutURL, &result.CreatedAt, &result.PaidAt, + &result.StripeCustomerID, &result.StripePaymentIntentID, &result.StripeChargeID, &result.StripeInvoiceID, + &result.InvoiceURL, &result.InvoicePDFURL, &result.ReceiptURL, &result.RefundedMicros, &result.DisputedMicros, + &result.ReconciliationStatus, &result.ReconciledAt, &result.ReconciliationError) if errors.Is(err, pgx.ErrNoRows) { return TopUpOrder{}, ErrTopUpOrderNotFound } diff --git a/internal/billing/operations.go b/internal/billing/operations.go new file mode 100644 index 0000000..a6dc653 --- /dev/null +++ b/internal/billing/operations.go @@ -0,0 +1,881 @@ +package billing + +import ( + "context" + "encoding/csv" + "encoding/json" + "errors" + "fmt" + "io" + "strconv" + "strings" + "time" + + "github.com/jackc/pgx/v5" + "github.com/stripe/stripe-go/v86" +) + +func (s *Service) CreatePortalSession(ctx context.Context, tenantID string) (PortalResult, error) { + if !s.stripeEnabled || s.stripeClient == nil { + return PortalResult{}, ErrStripeDisabled + } + var customerID string + if err := s.db.QueryRow(ctx, `SELECT stripe_customer_id FROM stripe_customers WHERE tenant_id=$1`, tenantID).Scan(&customerID); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return PortalResult{}, errors.New("no Stripe customer exists for this account") + } + return PortalResult{}, err + } + session, err := s.stripeClient.V1BillingPortalSessions.Create(ctx, &stripe.BillingPortalSessionCreateParams{ + Customer: stripe.String(customerID), ReturnURL: stripe.String(s.stripePortalReturnURL), + }) + if err != nil { + return PortalResult{}, fmt.Errorf("create Stripe customer portal session: %w", err) + } + if session.URL == "" { + return PortalResult{}, errors.New("Stripe returned an incomplete portal session") + } + return PortalResult{URL: session.URL}, nil +} + +func (s *Service) RetryCheckout(ctx context.Context, tenantID, orderID, email string) (CheckoutResult, error) { + var amount int64 + var status string + if err := s.db.QueryRow(ctx, `SELECT amount_minor,status FROM topup_orders WHERE id=$1 AND tenant_id=$2`, orderID, tenantID).Scan(&amount, &status); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return CheckoutResult{}, ErrTopUpOrderNotFound + } + return CheckoutResult{}, err + } + if status != "failed" && status != "expired" { + return CheckoutResult{}, errors.New("only failed or expired top-ups can be retried") + } + return s.CreateCheckout(ctx, CheckoutInput{TenantID: tenantID, AmountMinor: amount, CustomerEmail: email}) +} + +// ResolveMissingTopUp closes an uncredited local order only after reconciliation +// proved that its Checkout Session does not exist in the configured Stripe account. +func (s *Service) ResolveMissingTopUp(ctx context.Context, tenantID, orderID string, input ResolveMissingTopUpInput, actor ResolutionActor) (TopUpOrder, error) { + reason := normalizeDescription(input.Reason) + if reason == "" { + return TopUpOrder{}, fmt.Errorf("%w: a resolution reason is required", ErrCannotResolveTopUp) + } + tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}) + if err != nil { + return TopUpOrder{}, err + } + defer tx.Rollback(ctx) + var status, reconciliationStatus, paymentIntentID string + if err := tx.QueryRow(ctx, `SELECT status,reconciliation_status,COALESCE(stripe_payment_intent_id,'') + FROM topup_orders WHERE id=$1 AND tenant_id=$2 FOR UPDATE`, orderID, tenantID). + Scan(&status, &reconciliationStatus, &paymentIntentID); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return TopUpOrder{}, ErrTopUpOrderNotFound + } + return TopUpOrder{}, err + } + if status != "pending" || reconciliationStatus != "missing" || paymentIntentID != "" { + return TopUpOrder{}, fmt.Errorf("%w: only an uncredited pending order confirmed missing by reconciliation can be resolved", ErrCannotResolveTopUp) + } + var credits int64 + if err := tx.QueryRow(ctx, `SELECT count(*) FROM billing_ledger + WHERE source_type='stripe_checkout' AND source_id=(SELECT stripe_session_id FROM topup_orders WHERE id=$1)`, orderID).Scan(&credits); err != nil { + return TopUpOrder{}, err + } + if credits != 0 { + return TopUpOrder{}, fmt.Errorf("%w: a credited top-up cannot be resolved as missing", ErrCannotResolveTopUp) + } + if actor.Type != "console_user" && actor.Type != "bootstrap" && actor.Type != "maintenance" { + return TopUpOrder{}, fmt.Errorf("%w: a valid resolution actor is required", ErrCannotResolveTopUp) + } + if _, err := tx.Exec(ctx, `INSERT INTO billing_reconciliation_resolutions + (topup_order_id,tenant_id,actor_id,actor_type,reason) VALUES ($1,$2,$3,$4,$5)`, + orderID, tenantID, actor.ID, actor.Type, reason); err != nil { + return TopUpOrder{}, err + } + if _, err := tx.Exec(ctx, `UPDATE topup_orders SET status='failed',reconciliation_status='resolved', + reconciled_at=now(),reconciliation_error=$2 WHERE id=$1`, orderID, reason); err != nil { + return TopUpOrder{}, err + } + if err := tx.Commit(ctx); err != nil { + return TopUpOrder{}, err + } + return s.GetTopUpOrder(ctx, tenantID, orderID) +} + +// ReverseMissingTopUpCredit preserves the original credit and adds an equal +// negative ledger entry when the configured Stripe account cannot prove the +// payment. It refuses to consume funds reserved for in-flight requests. +func (s *Service) ReverseMissingTopUpCredit(ctx context.Context, tenantID, orderID string, input ResolveMissingTopUpInput, actor ResolutionActor) (TopUpOrder, error) { + reason := normalizeDescription(input.Reason) + if reason == "" || (actor.Type != "console_user" && actor.Type != "bootstrap" && actor.Type != "maintenance") { + return TopUpOrder{}, fmt.Errorf("%w: a reason and valid resolution actor are required", ErrCannotResolveTopUp) + } + tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}) + if err != nil { + return TopUpOrder{}, err + } + defer tx.Rollback(ctx) + var status, reconciliationStatus, paymentIntentID, currency, sessionID string + var amount int64 + if err := tx.QueryRow(ctx, `SELECT status,reconciliation_status,COALESCE(stripe_payment_intent_id,''), + currency,amount_micros,COALESCE(stripe_session_id,'') FROM topup_orders + WHERE id=$1 AND tenant_id=$2 FOR UPDATE`, orderID, tenantID). + Scan(&status, &reconciliationStatus, &paymentIntentID, ¤cy, &amount, &sessionID); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return TopUpOrder{}, ErrTopUpOrderNotFound + } + return TopUpOrder{}, err + } + if status != "paid" || reconciliationStatus != "missing" || paymentIntentID != "" || sessionID == "" { + return TopUpOrder{}, fmt.Errorf("%w: only a paid credit confirmed missing with no PaymentIntent can be reversed", ErrCannotResolveTopUp) + } + var originalCredits int64 + if err := tx.QueryRow(ctx, `SELECT COALESCE(sum(amount_micros),0) FROM billing_ledger + WHERE tenant_id=$1 AND source_type='stripe_checkout' AND source_id=$2`, tenantID, sessionID).Scan(&originalCredits); err != nil { + return TopUpOrder{}, err + } + if originalCredits != amount { + return TopUpOrder{}, fmt.Errorf("%w: original Stripe credit does not match the order", ErrCannotResolveTopUp) + } + var balance, reserved int64 + if err := tx.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1 FOR UPDATE`, tenantID).Scan(&balance, &reserved); err != nil { + return TopUpOrder{}, err + } + if balance-reserved < amount { + return TopUpOrder{}, fmt.Errorf("%w: available balance is insufficient to reverse the orphaned credit", ErrCannotResolveTopUp) + } + newBalance := balance - amount + if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET balance_micros=$2,updated_at=now() WHERE tenant_id=$1`, tenantID, newBalance); err != nil { + return TopUpOrder{}, err + } + if _, err := tx.Exec(ctx, `INSERT INTO billing_ledger + (tenant_id,currency,amount_micros,balance_after_micros,kind,source_type,source_id,description) + VALUES ($1,$2,$3,$4,'adjustment','stripe_reconciliation',$5,$6)`, + tenantID, currency, -amount, newBalance, orderID, reason); err != nil { + return TopUpOrder{}, err + } + if _, err := tx.Exec(ctx, `INSERT INTO billing_reconciliation_resolutions + (topup_order_id,tenant_id,actor_id,actor_type,reason) VALUES ($1,$2,$3,$4,$5)`, + orderID, tenantID, actor.ID, actor.Type, reason); err != nil { + return TopUpOrder{}, err + } + if _, err := tx.Exec(ctx, `UPDATE topup_orders SET status='reversed',reconciliation_status='resolved', + reconciled_at=now(),reconciliation_error=$2 WHERE id=$1`, orderID, reason); err != nil { + return TopUpOrder{}, err + } + if err := tx.Commit(ctx); err != nil { + return TopUpOrder{}, err + } + return s.GetTopUpOrder(ctx, tenantID, orderID) +} + +func (s *Service) CreateRefund(ctx context.Context, tenantID, orderID string, input RefundInput) (Refund, error) { + if !s.stripeEnabled { + return Refund{}, ErrStripeDisabled + } + input.Reason = strings.TrimSpace(input.Reason) + if input.Reason == "" { + input.Reason = "requested_by_customer" + } + if input.Reason != "requested_by_customer" && input.Reason != "duplicate" && input.Reason != "fraudulent" { + return Refund{}, errors.New("refund reason must be requested_by_customer, duplicate, or fraudulent") + } + tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}) + if err != nil { + return Refund{}, err + } + defer tx.Rollback(ctx) + var orderTenant, currency, status, paymentIntentID string + var orderAmount, refunded int64 + if err := tx.QueryRow(ctx, `SELECT tenant_id::text,currency,status,COALESCE(stripe_payment_intent_id,''),amount_micros,refunded_micros + FROM topup_orders WHERE id=$1 AND tenant_id=$2 FOR UPDATE`, orderID, tenantID).Scan(&orderTenant, ¤cy, &status, &paymentIntentID, &orderAmount, &refunded); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return Refund{}, ErrTopUpOrderNotFound + } + return Refund{}, err + } + if status != "paid" && status != "partially_refunded" { + return Refund{}, errors.New("only a paid top-up can be refunded") + } + if paymentIntentID == "" { + return Refund{}, errors.New("top-up has no Stripe PaymentIntent") + } + amountMicros, err := minorToMicros(currency, input.AmountMinor) + if err != nil { + return Refund{}, ErrInvalidAmount + } + var pendingRefunds int64 + if err := tx.QueryRow(ctx, `SELECT COALESCE(sum(amount_micros),0) FROM stripe_refunds + WHERE topup_order_id=$1 AND status IN ('queued','submitting','pending','requires_action','succeeded')`, orderID).Scan(&pendingRefunds); err != nil { + return Refund{}, err + } + // Legacy successful refunds are already reflected in refunded_micros. New + // rows are included in the sum, so use the larger value without double-counting. + committedRefunds := refunded + if pendingRefunds > committedRefunds { + committedRefunds = pendingRefunds + } + if amountMicros > orderAmount-committedRefunds { + return Refund{}, ErrInvalidAmount + } + var balance, held int64 + if err := tx.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1 FOR UPDATE`, tenantID).Scan(&balance, &held); err != nil { + return Refund{}, err + } + available := balance - held + if available < 0 { + available = 0 + } + refundHold := amountMicros + if refundHold > available { + refundHold = available + } + if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET reserved_micros=reserved_micros+$2,updated_at=now() WHERE tenant_id=$1`, tenantID, refundHold); err != nil { + return Refund{}, err + } + var result Refund + err = tx.QueryRow(ctx, `INSERT INTO stripe_refunds (tenant_id,topup_order_id,amount_minor,amount_micros,held_micros,currency,reason) + VALUES ($1,$2,$3,$4,$5,$6,$7) RETURNING id::text,tenant_id::text,topup_order_id::text,amount_minor,amount_micros,currency,reason,status,last_error,created_at,completed_at`, + tenantID, orderID, input.AmountMinor, amountMicros, refundHold, currency, input.Reason).Scan(&result.ID, &result.TenantID, &result.TopUpOrderID, + &result.AmountMinor, &result.AmountMicros, &result.Currency, &result.Reason, &result.Status, &result.LastError, &result.CreatedAt, &result.CompletedAt) + if err != nil { + return Refund{}, err + } + if err := tx.Commit(ctx); err != nil { + return Refund{}, err + } + return result, nil +} + +func (s *Service) ListRefunds(ctx context.Context, tenantID string, limit int) ([]Refund, error) { + if limit < 1 || limit > 500 { + limit = 100 + } + query := `SELECT id::text,tenant_id::text,topup_order_id::text,COALESCE(stripe_refund_id,''),amount_minor,amount_micros, + currency,reason,status,last_error,created_at,completed_at FROM stripe_refunds` + args := []any{} + if tenantID != "" { + query += ` WHERE tenant_id=$1 ORDER BY created_at DESC LIMIT $2` + args = []any{tenantID, limit} + } else { + query += ` ORDER BY created_at DESC LIMIT $1` + args = []any{limit} + } + rows, err := s.db.Query(ctx, query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + result := make([]Refund, 0) + for rows.Next() { + var item Refund + if err := rows.Scan(&item.ID, &item.TenantID, &item.TopUpOrderID, &item.StripeRefundID, &item.AmountMinor, &item.AmountMicros, &item.Currency, &item.Reason, &item.Status, &item.LastError, &item.CreatedAt, &item.CompletedAt); err != nil { + return nil, err + } + result = append(result, item) + } + return result, rows.Err() +} + +func (s *Service) ListDisputes(ctx context.Context, tenantID string, limit int) ([]PaymentDispute, error) { + if limit < 1 || limit > 500 { + limit = 100 + } + query := `SELECT stripe_dispute_id,COALESCE(tenant_id::text,''),COALESCE(topup_order_id::text,''), + amount_minor,amount_micros,currency,status,reason,debited_micros,uncollected_micros,due_by,updated_at FROM stripe_disputes` + args := []any{} + if tenantID != "" { + query += ` WHERE tenant_id=$1 ORDER BY updated_at DESC LIMIT $2` + args = []any{tenantID, limit} + } else { + query += ` ORDER BY updated_at DESC LIMIT $1` + args = []any{limit} + } + rows, err := s.db.Query(ctx, query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + result := []PaymentDispute{} + for rows.Next() { + var item PaymentDispute + if err := rows.Scan(&item.ID, &item.TenantID, &item.TopUpOrderID, &item.AmountMinor, &item.AmountMicros, &item.Currency, &item.Status, &item.Reason, &item.DebitedMicros, &item.UncollectedMicros, &item.DueBy, &item.UpdatedAt); err != nil { + return nil, err + } + result = append(result, item) + } + return result, rows.Err() +} + +func (s *Service) ListInvoices(ctx context.Context, tenantID string, limit int) ([]Invoice, error) { + if limit < 1 || limit > 500 { + limit = 100 + } + query := `SELECT stripe_invoice_id,COALESCE(tenant_id::text,''),COALESCE(topup_order_id::text,''),status,currency, + amount_due_minor,amount_paid_minor,attempt_count,next_payment_attempt,hosted_invoice_url,invoice_pdf_url,last_failure,updated_at FROM stripe_invoices` + args := []any{} + if tenantID != "" { + query += ` WHERE tenant_id=$1 ORDER BY updated_at DESC LIMIT $2` + args = []any{tenantID, limit} + } else { + query += ` ORDER BY updated_at DESC LIMIT $1` + args = []any{limit} + } + rows, err := s.db.Query(ctx, query, args...) + if err != nil { + return nil, err + } + defer rows.Close() + result := []Invoice{} + for rows.Next() { + var item Invoice + if err := rows.Scan(&item.ID, &item.TenantID, &item.TopUpOrderID, &item.Status, &item.Currency, &item.AmountDueMinor, &item.AmountPaidMinor, &item.AttemptCount, &item.NextPaymentAttempt, &item.HostedInvoiceURL, &item.InvoicePDFURL, &item.LastFailure, &item.UpdatedAt); err != nil { + return nil, err + } + result = append(result, item) + } + return result, rows.Err() +} + +func (s *Service) RunStripeOperations(ctx context.Context) { + if !s.stripeEnabled || s.stripeClient == nil { + return + } + ticker := time.NewTicker(2 * time.Second) + reconcile := time.NewTicker(30 * time.Minute) + metrics := time.NewTicker(15 * time.Second) + defer ticker.Stop() + defer reconcile.Stop() + defer metrics.Stop() + initialCtx, cancel := context.WithTimeout(ctx, 2*time.Minute) + _, _ = s.Reconcile(initialCtx, 100) + cancel() + s.refreshOperationalMetrics(ctx) + for { + for i := 0; i < 8; i++ { + ok, _ := s.processRefundOperation(ctx) + if !ok { + break + } + } + _, _ = s.pollPendingRefund(ctx) + select { + case <-ctx.Done(): + return + case <-ticker.C: + case <-reconcile.C: + _, _ = s.Reconcile(ctx, 100) + case <-metrics.C: + s.refreshOperationalMetrics(ctx) + } + } +} + +func (s *Service) refreshOperationalMetrics(ctx context.Context) { + if s.metrics == nil { + return + } + status, err := s.OperationalStatus(ctx) + if err == nil { + s.metrics.SetStripeOperations(status) + } +} + +func (s *Service) OperationalStatus(ctx context.Context) (OperationalStatus, error) { + status := OperationalStatus{StripeEnabled: s.stripeEnabled, ReconciliationStatus: "disabled"} + if err := s.db.QueryRow(ctx, `SELECT + count(*) FILTER (WHERE status IN ('queued','submitting','pending','requires_action')), + min(created_at) FILTER (WHERE status IN ('queued','submitting','pending','requires_action')), + COALESCE((SELECT sum(uncollected_micros) FROM stripe_refunds),0)+ + COALESCE((SELECT sum(uncollected_micros) FROM stripe_disputes),0)+ + COALESCE((SELECT sum(uncollected_micros) FROM usage_events),0) + FROM stripe_refunds`).Scan(&status.RefundBacklog, &status.OldestRefund, &status.UncollectedMicros); err != nil { + return status, err + } + if err := s.db.QueryRow(ctx, `SELECT count(*),min(created_at) FROM stripe_webhook_events + WHERE processed_at IS NULL`).Scan(&status.UnprocessedWebhooks, &status.OldestUnprocessedWebhook); err != nil { + return status, err + } + if err := s.db.QueryRow(ctx, `SELECT count(*) FROM usage_events WHERE metering_status='missing'`).Scan(&status.UnmeteredSuccesses); err != nil { + return status, err + } + if !s.stripeEnabled { + return status, nil + } + err := s.db.QueryRow(ctx, `SELECT status,mismatch_count,completed_at,error FROM billing_reconciliation_runs + WHERE status<>'running' ORDER BY started_at DESC LIMIT 1`).Scan(&status.ReconciliationStatus, + &status.ReconciliationMismatches, &status.ReconciliationCompletedAt, &status.ReconciliationError) + if errors.Is(err, pgx.ErrNoRows) { + status.ReconciliationStatus = "never_run" + return status, nil + } + return status, err +} + +func (s *Service) processRefundOperation(ctx context.Context) (bool, error) { + tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}) + if err != nil { + return false, err + } + defer tx.Rollback(ctx) + var id, orderID, paymentIntentID, reason string + var amount int64 + err = tx.QueryRow(ctx, `WITH selected AS ( + SELECT r.id FROM stripe_refunds r WHERE r.available_at<=now() AND + (r.status='queued' OR (r.status='submitting' AND r.updated_at<now()-interval '5 minutes')) + ORDER BY r.available_at,r.created_at FOR UPDATE SKIP LOCKED LIMIT 1) + UPDATE stripe_refunds r SET status='submitting',attempts=attempts+1,updated_at=now() + FROM selected,topup_orders o WHERE r.id=selected.id AND o.id=r.topup_order_id + RETURNING r.id::text,r.topup_order_id::text,o.stripe_payment_intent_id,r.amount_minor,r.reason`).Scan(&id, &orderID, &paymentIntentID, &amount, &reason) + if errors.Is(err, pgx.ErrNoRows) { + return false, tx.Commit(ctx) + } + if err != nil { + return false, err + } + if err := tx.Commit(ctx); err != nil { + return false, err + } + params := &stripe.RefundCreateParams{Amount: stripe.Int64(amount), PaymentIntent: stripe.String(paymentIntentID), Reason: stripe.String(reason), + Metadata: map[string]string{"aigw_refund_id": id, "aigw_topup_order_id": orderID}} + params.SetIdempotencyKey("aigw_refund_" + id) + refund, err := s.stripeClient.V1Refunds.Create(ctx, params) + if err != nil { + message := err.Error() + if len(message) > 1000 { + message = message[:1000] + } + _, _ = s.db.Exec(ctx, `UPDATE stripe_refunds SET status='queued',last_error=$2, + available_at=now()+make_interval(secs=>LEAST(1800,power(2,LEAST(attempts,10))::int)),updated_at=now() WHERE id=$1`, id, message) + return true, err + } + return true, s.applyRefund(ctx, refund) +} + +func (s *Service) pollPendingRefund(ctx context.Context) (bool, error) { + var id string + err := s.db.QueryRow(ctx, `UPDATE stripe_refunds SET available_at=now()+interval '1 minute',updated_at=now() + WHERE id=(SELECT id FROM stripe_refunds WHERE status IN ('pending','requires_action') + AND stripe_refund_id IS NOT NULL AND available_at<=now() ORDER BY available_at LIMIT 1 FOR UPDATE SKIP LOCKED) + RETURNING stripe_refund_id`).Scan(&id) + if errors.Is(err, pgx.ErrNoRows) { + return false, nil + } + if err != nil { + return false, err + } + refund, err := s.stripeClient.V1Refunds.Retrieve(ctx, id, &stripe.RefundRetrieveParams{}) + if err != nil { + return true, err + } + return true, s.applyRefund(ctx, refund) +} + +func (s *Service) applyRefund(ctx context.Context, refund *stripe.Refund) error { + if refund == nil || refund.ID == "" { + return ErrInvalidAmount + } + tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}) + if err != nil { + return err + } + defer tx.Rollback(ctx) + if err := s.applyRefundTx(ctx, tx, refund); err != nil { + return err + } + return tx.Commit(ctx) +} + +func (s *Service) applyRefundTx(ctx context.Context, tx pgx.Tx, refund *stripe.Refund) error { + if refund == nil || refund.ID == "" { + return ErrInvalidAmount + } + localID := refund.Metadata["aigw_refund_id"] + var id, tenantID, orderID, currentStatus string + var amountMicros, held int64 + query := `SELECT id::text,tenant_id::text,topup_order_id::text,amount_micros,held_micros,status FROM stripe_refunds WHERE ` + arg := refund.ID + if localID != "" { + query += `id=$1 FOR UPDATE` + arg = localID + } else { + query += `stripe_refund_id=$1 FOR UPDATE` + } + err := tx.QueryRow(ctx, query, arg).Scan(&id, &tenantID, &orderID, &amountMicros, &held, ¤tStatus) + if errors.Is(err, pgx.ErrNoRows) { + paymentIntentID := "" + if refund.PaymentIntent != nil { + paymentIntentID = refund.PaymentIntent.ID + } + if paymentIntentID == "" { + return ErrInvalidAmount + } + var currency string + if err := tx.QueryRow(ctx, `SELECT id::text,tenant_id::text,currency FROM topup_orders WHERE stripe_payment_intent_id=$1 FOR UPDATE`, paymentIntentID).Scan(&orderID, &tenantID, ¤cy); err != nil { + return err + } + amountMicros, err = minorToMicros(currency, refund.Amount) + if err != nil { + return err + } + err = tx.QueryRow(ctx, `INSERT INTO stripe_refunds (tenant_id,topup_order_id,stripe_refund_id,amount_minor,amount_micros,currency,reason,status) + VALUES ($1,$2,$3,$4,$5,$6,$7,$8) RETURNING id::text`, tenantID, orderID, refund.ID, refund.Amount, amountMicros, currency, string(refund.Reason), string(refund.Status)).Scan(&id) + if err != nil { + return err + } + currentStatus = "" + } else if err != nil { + return err + } + if _, err := tx.Exec(ctx, `UPDATE stripe_refunds SET stripe_refund_id=$2,status=$3,last_error=$4,updated_at=now(), + completed_at=CASE WHEN $3 IN ('succeeded','failed','canceled') THEN now() ELSE completed_at END WHERE id=$1`, id, refund.ID, string(refund.Status), string(refund.FailureReason)); err != nil { + return err + } + if refund.Status == stripe.RefundStatusSucceeded && currentStatus != "succeeded" { + var balance, reserved int64 + if err := tx.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1 FOR UPDATE`, tenantID).Scan(&balance, &reserved); err != nil { + return err + } + if held > 0 { + if held > reserved { + return errors.New("refund hold invariant violated") + } + reserved -= held + } + debit := amountMicros + if debit > balance-reserved { + debit = balance - reserved + } + if debit < 0 { + debit = 0 + } + newBalance := balance - debit + if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET balance_micros=$2,reserved_micros=$3,updated_at=now() WHERE tenant_id=$1`, tenantID, newBalance, reserved); err != nil { + return err + } + if _, err := tx.Exec(ctx, `UPDATE stripe_refunds SET held_micros=0,uncollected_micros=$2 WHERE id=$1`, id, amountMicros-debit); err != nil { + return err + } + if _, err := tx.Exec(ctx, `INSERT INTO billing_ledger (tenant_id,currency,amount_micros,balance_after_micros,kind,source_type,source_id,description) + SELECT $1,currency,$2,$3,'refund','stripe_refund',$4,'Stripe top-up refund' FROM topup_orders WHERE id=$5 + ON CONFLICT (source_type,source_id) DO NOTHING`, tenantID, -debit, newBalance, refund.ID, orderID); err != nil { + return err + } + if _, err := tx.Exec(ctx, `UPDATE topup_orders SET refunded_micros=LEAST(amount_micros,refunded_micros+$2), + status=CASE WHEN refunded_micros+$2>=amount_micros THEN 'refunded' ELSE 'partially_refunded' END WHERE id=$1`, orderID, amountMicros); err != nil { + return err + } + } else if (refund.Status == stripe.RefundStatusFailed || refund.Status == stripe.RefundStatusCanceled) && held > 0 { + if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET reserved_micros=reserved_micros-$2,updated_at=now() WHERE tenant_id=$1`, tenantID, held); err != nil { + return err + } + if _, err := tx.Exec(ctx, `UPDATE stripe_refunds SET held_micros=0 WHERE id=$1`, id); err != nil { + return err + } + } + return nil +} + +func (s *Service) applyDisputeTx(ctx context.Context, tx pgx.Tx, dispute *stripe.Dispute, eventType stripe.EventType) error { + paymentIntentID := "" + if dispute.PaymentIntent != nil { + paymentIntentID = dispute.PaymentIntent.ID + } + if paymentIntentID == "" && dispute.Charge != nil && dispute.Charge.PaymentIntent != nil { + paymentIntentID = dispute.Charge.PaymentIntent.ID + } + if paymentIntentID == "" { + return ErrInvalidAmount + } + var orderID, tenantID, currency string + if err := tx.QueryRow(ctx, `SELECT id::text,tenant_id::text,currency FROM topup_orders WHERE stripe_payment_intent_id=$1 FOR UPDATE`, paymentIntentID).Scan(&orderID, &tenantID, ¤cy); err != nil { + return err + } + amountMicros, err := minorToMicros(currency, dispute.Amount) + if err != nil { + return err + } + dueBy := (*time.Time)(nil) + if dispute.EvidenceDetails != nil && dispute.EvidenceDetails.DueBy > 0 { + value := time.Unix(dispute.EvidenceDetails.DueBy, 0).UTC() + dueBy = &value + } + var previousDebited int64 + queryErr := tx.QueryRow(ctx, `SELECT debited_micros FROM stripe_disputes WHERE stripe_dispute_id=$1 FOR UPDATE`, dispute.ID).Scan(&previousDebited) + if queryErr != nil && !errors.Is(queryErr, pgx.ErrNoRows) { + return queryErr + } + if _, err := tx.Exec(ctx, `INSERT INTO stripe_disputes (stripe_dispute_id,tenant_id,topup_order_id,stripe_payment_intent_id, + amount_minor,amount_micros,currency,status,reason,due_by) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10) + ON CONFLICT (stripe_dispute_id) DO UPDATE SET status=EXCLUDED.status,reason=EXCLUDED.reason, + due_by=EXCLUDED.due_by,updated_at=now(),closed_at=CASE WHEN EXCLUDED.status IN ('won','lost') THEN now() ELSE stripe_disputes.closed_at END`, + dispute.ID, tenantID, orderID, paymentIntentID, dispute.Amount, amountMicros, currency, string(dispute.Status), string(dispute.Reason), dueBy); err != nil { + return err + } + shouldDebit := eventType == stripe.EventTypeChargeDisputeCreated || eventType == stripe.EventTypeChargeDisputeFundsWithdrawn + shouldReverse := eventType == stripe.EventTypeChargeDisputeFundsReinstated || dispute.Status == stripe.DisputeStatusWon + if shouldDebit && previousDebited == 0 { + var balance, reserved int64 + if err := tx.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1 FOR UPDATE`, tenantID).Scan(&balance, &reserved); err != nil { + return err + } + debit := amountMicros + if debit > balance-reserved { + debit = balance - reserved + } + if debit < 0 { + debit = 0 + } + newBalance := balance - debit + if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET balance_micros=$2,updated_at=now() WHERE tenant_id=$1`, tenantID, newBalance); err != nil { + return err + } + if _, err := tx.Exec(ctx, `UPDATE stripe_disputes SET debited_micros=$2,uncollected_micros=$3 WHERE stripe_dispute_id=$1`, dispute.ID, debit, amountMicros-debit); err != nil { + return err + } + if _, err := tx.Exec(ctx, `INSERT INTO billing_ledger (tenant_id,currency,amount_micros,balance_after_micros,kind,source_type,source_id,description) + VALUES ($1,$2,$3,$4,'dispute','stripe_dispute',$5,'Stripe payment dispute') ON CONFLICT DO NOTHING`, tenantID, currency, -debit, newBalance, dispute.ID); err != nil { + return err + } + if _, err := tx.Exec(ctx, `UPDATE topup_orders SET disputed_micros=GREATEST(disputed_micros,$2),status='disputed' WHERE id=$1`, orderID, amountMicros); err != nil { + return err + } + } else if shouldReverse && previousDebited > 0 { + var balance int64 + if err := tx.QueryRow(ctx, `SELECT balance_micros FROM tenant_wallets WHERE tenant_id=$1 FOR UPDATE`, tenantID).Scan(&balance); err != nil { + return err + } + newBalance := balance + previousDebited + if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET balance_micros=$2,updated_at=now() WHERE tenant_id=$1`, tenantID, newBalance); err != nil { + return err + } + if _, err := tx.Exec(ctx, `UPDATE stripe_disputes SET debited_micros=0,uncollected_micros=0 WHERE stripe_dispute_id=$1`, dispute.ID); err != nil { + return err + } + if _, err := tx.Exec(ctx, `INSERT INTO billing_ledger (tenant_id,currency,amount_micros,balance_after_micros,kind,source_type,source_id,description) + VALUES ($1,$2,$3,$4,'dispute_reversal','stripe_dispute_reversal',$5,'Stripe dispute funds reinstated') ON CONFLICT DO NOTHING`, tenantID, currency, previousDebited, newBalance, dispute.ID); err != nil { + return err + } + if _, err := tx.Exec(ctx, `UPDATE topup_orders SET disputed_micros=0,status=CASE WHEN refunded_micros=0 THEN 'paid' WHEN refunded_micros<amount_micros THEN 'partially_refunded' ELSE 'refunded' END WHERE id=$1`, orderID); err != nil { + return err + } + } + return nil +} + +func (s *Service) applyInvoiceTx(ctx context.Context, tx pgx.Tx, invoice *stripe.Invoice) error { + orderID := invoice.Metadata["aigw_topup_order_id"] + tenantID := invoice.Metadata["aigw_tenant_id"] + customerID := "" + if invoice.Customer != nil { + customerID = invoice.Customer.ID + } + if tenantID == "" && customerID != "" { + _ = tx.QueryRow(ctx, `SELECT tenant_id::text FROM stripe_customers WHERE stripe_customer_id=$1`, customerID).Scan(&tenantID) + } + if orderID == "" { + _ = tx.QueryRow(ctx, `SELECT id::text FROM topup_orders WHERE stripe_invoice_id=$1`, invoice.ID).Scan(&orderID) + } + failure := "" + if invoice.LastFinalizationError != nil { + failure = invoice.LastFinalizationError.Msg + } + var nextAttempt *time.Time + if invoice.NextPaymentAttempt > 0 { + value := time.Unix(invoice.NextPaymentAttempt, 0).UTC() + nextAttempt = &value + } + if _, err := tx.Exec(ctx, `INSERT INTO stripe_invoices (stripe_invoice_id,tenant_id,topup_order_id,stripe_customer_id,status,currency, + amount_due_minor,amount_paid_minor,attempt_count,next_payment_attempt,hosted_invoice_url,invoice_pdf_url,last_failure) + VALUES ($1,NULLIF($2,'')::uuid,NULLIF($3,'')::uuid,NULLIF($4,''),$5,$6,$7,$8,$9,$10,$11,$12,$13) + ON CONFLICT (stripe_invoice_id) DO UPDATE SET status=EXCLUDED.status,amount_due_minor=EXCLUDED.amount_due_minor, + amount_paid_minor=EXCLUDED.amount_paid_minor,attempt_count=EXCLUDED.attempt_count,next_payment_attempt=EXCLUDED.next_payment_attempt, + hosted_invoice_url=EXCLUDED.hosted_invoice_url,invoice_pdf_url=EXCLUDED.invoice_pdf_url,last_failure=EXCLUDED.last_failure,updated_at=now()`, + invoice.ID, tenantID, orderID, customerID, string(invoice.Status), string(invoice.Currency), invoice.AmountDue, invoice.AmountPaid, + invoice.AttemptCount, nextAttempt, invoice.HostedInvoiceURL, invoice.InvoicePDF, failure); err != nil { + return err + } + if orderID != "" { + _, err := tx.Exec(ctx, `UPDATE topup_orders SET stripe_invoice_id=$2,invoice_url=COALESCE(NULLIF($3,''),invoice_url), + invoice_pdf_url=COALESCE(NULLIF($4,''),invoice_pdf_url) WHERE id=$1`, orderID, invoice.ID, invoice.HostedInvoiceURL, invoice.InvoicePDF) + return err + } + return nil +} + +func (s *Service) Reconcile(ctx context.Context, limit int) (ReconciliationResult, error) { + if !s.stripeEnabled || s.stripeClient == nil { + return ReconciliationResult{}, ErrStripeDisabled + } + if limit < 1 || limit > 500 { + limit = 100 + } + var result ReconciliationResult + if err := s.db.QueryRow(ctx, `INSERT INTO billing_reconciliation_runs (status) VALUES ('running') RETURNING id::text`).Scan(&result.ID); err != nil { + return result, err + } + rows, err := s.db.Query(ctx, `SELECT id::text,stripe_session_id,status,amount_minor,currency,reconciliation_status FROM topup_orders + WHERE stripe_session_id IS NOT NULL ORDER BY created_at DESC LIMIT $1`, limit) + if err != nil { + return s.failReconciliation(ctx, result, err) + } + type order struct { + id, session, status, currency, reconciliationStatus string + amount int64 + } + var orders []order + for rows.Next() { + var item order + if err := rows.Scan(&item.id, &item.session, &item.status, &item.amount, &item.currency, &item.reconciliationStatus); err != nil { + rows.Close() + return s.failReconciliation(ctx, result, err) + } + orders = append(orders, item) + } + rows.Close() + for _, item := range orders { + session, retrieveErr := s.stripeClient.V1CheckoutSessions.Retrieve(ctx, item.session, &stripe.CheckoutSessionRetrieveParams{}) + result.CheckedOrders++ + if retrieveErr != nil { + if stripeResourceMissing(retrieveErr) && item.reconciliationStatus == "resolved" && (item.status == "failed" || item.status == "reversed") { + continue + } + message := truncateError(retrieveErr) + typeName := "stripe_retrieve_failed" + state := "mismatch" + if stripeResourceMissing(retrieveErr) { + typeName = "stripe_session_missing" + state = "missing" + } + if err := s.updateOrderReconciliation(ctx, item.id, state, message); err != nil { + return s.failReconciliation(ctx, result, fmt.Errorf("record reconciliation retrieval failure for order %s: %w", item.id, err)) + } + result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": typeName, "error": message}) + continue + } + expected := item.status == "paid" || item.status == "partially_refunded" || item.status == "refunded" || item.status == "disputed" + stripePaid := session.PaymentStatus == stripe.CheckoutSessionPaymentStatusPaid + if session.AmountTotal != item.amount || string(session.Currency) != item.currency { + result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "checkout_mismatch", "local_status": item.status, "stripe_payment_status": session.PaymentStatus, "local_amount": item.amount, "stripe_amount": session.AmountTotal}) + if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", "amount or currency mismatch"); err != nil { + return s.failReconciliation(ctx, result, fmt.Errorf("record checkout mismatch for order %s: %w", item.id, err)) + } + continue + } + if stripePaid && !expected { + raw, marshalErr := json.Marshal(session) + if marshalErr != nil { + return s.failReconciliation(ctx, result, fmt.Errorf("encode Stripe Checkout Session for order %s: %w", item.id, marshalErr)) + } + if repairErr := s.processStripeEvent(ctx, stripe.Event{ID: "reconcile_" + session.ID + "_paid", Type: stripe.EventTypeCheckoutSessionCompleted, Data: &stripe.EventData{Raw: raw}}); repairErr != nil { + message := truncateError(repairErr) + result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "checkout_repair_failed", "error": message}) + if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", message); err != nil { + return s.failReconciliation(ctx, result, fmt.Errorf("record checkout repair failure for order %s: %w", item.id, err)) + } + continue + } + result.Repairs = append(result.Repairs, map[string]any{"order_id": item.id, "type": "credited_paid_checkout"}) + if err := s.updateOrderReconciliation(ctx, item.id, "repaired", ""); err != nil { + return s.failReconciliation(ctx, result, fmt.Errorf("record checkout repair for order %s: %w", item.id, err)) + } + continue + } + if expected != stripePaid { + result.Mismatches = append(result.Mismatches, map[string]any{"order_id": item.id, "type": "checkout_payment_state_mismatch", "local_status": item.status, "stripe_payment_status": session.PaymentStatus}) + if err := s.updateOrderReconciliation(ctx, item.id, "mismatch", "payment state mismatch"); err != nil { + return s.failReconciliation(ctx, result, fmt.Errorf("record payment-state mismatch for order %s: %w", item.id, err)) + } + continue + } + if err := s.updateOrderReconciliation(ctx, item.id, "ok", ""); err != nil { + return s.failReconciliation(ctx, result, fmt.Errorf("record clean reconciliation for order %s: %w", item.id, err)) + } + } + result.MismatchCount = int64(len(result.Mismatches)) + result.Status = "clean" + if result.MismatchCount > 0 { + result.Status = "mismatch" + } + report, marshalErr := json.Marshal(result.Mismatches) + if marshalErr != nil { + return s.failReconciliation(ctx, result, fmt.Errorf("encode reconciliation report: %w", marshalErr)) + } + _, err = s.db.Exec(ctx, `UPDATE billing_reconciliation_runs SET status=$2,checked_orders=$3,mismatch_count=$4,report=$5,completed_at=now() WHERE id=$1`, result.ID, result.Status, result.CheckedOrders, result.MismatchCount, report) + if err != nil { + return s.failReconciliation(ctx, result, fmt.Errorf("complete reconciliation run: %w", err)) + } + return result, nil +} + +func (s *Service) updateOrderReconciliation(ctx context.Context, orderID, status, message string) error { + command, err := s.db.Exec(ctx, `UPDATE topup_orders SET reconciliation_status=$2,reconciled_at=now(),reconciliation_error=$3 WHERE id=$1`, orderID, status, message) + if err != nil { + return err + } + if command.RowsAffected() != 1 { + return fmt.Errorf("expected one top-up order, updated %d", command.RowsAffected()) + } + return nil +} + +func stripeResourceMissing(err error) bool { + var stripeErr *stripe.Error + return errors.As(err, &stripeErr) && (stripeErr.Code == stripe.ErrorCodeResourceMissing || stripeErr.HTTPStatusCode == 404) +} + +func truncateError(err error) string { + if err == nil { + return "" + } + message := err.Error() + if len(message) > 1000 { + return message[:1000] + } + return message +} + +func (s *Service) failReconciliation(ctx context.Context, result ReconciliationResult, cause error) (ReconciliationResult, error) { + if _, err := s.db.Exec(ctx, `UPDATE billing_reconciliation_runs SET status='failed',error=$2,completed_at=now() WHERE id=$1`, result.ID, truncateError(cause)); err != nil { + return result, errors.Join(cause, fmt.Errorf("record failed reconciliation run: %w", err)) + } + return result, cause +} + +func (s *Service) WriteFinancialCSV(ctx context.Context, tenantID string, output io.Writer) error { + query := ` + SELECT id::text, tenant_id::text, COALESCE(project_id::text, ''), currency, amount_micros, + balance_after_micros, kind, source_type, source_id, description, created_at + FROM billing_ledger` + args := []any{} + if strings.TrimSpace(tenantID) != "" { + query += ` WHERE tenant_id=$1` + args = append(args, tenantID) + } + query += ` ORDER BY created_at, id` + rows, err := s.db.Query(ctx, query, args...) + if err != nil { + return err + } + defer rows.Close() + w := csv.NewWriter(output) + if err := w.Write([]string{"timestamp", "tenant_id", "project_id", "currency", "kind", "amount_micros", "balance_after_micros", "source_type", "source_id", "description"}); err != nil { + return err + } + for rows.Next() { + var e LedgerEntry + if err := rows.Scan(&e.ID, &e.TenantID, &e.ProjectID, &e.Currency, &e.AmountMicros, + &e.BalanceAfterMicros, &e.Kind, &e.SourceType, &e.SourceID, &e.Description, &e.CreatedAt); err != nil { + return err + } + if err := w.Write([]string{e.CreatedAt.UTC().Format(time.RFC3339Nano), e.TenantID, e.ProjectID, e.Currency, e.Kind, strconv.FormatInt(e.AmountMicros, 10), strconv.FormatInt(e.BalanceAfterMicros, 10), e.SourceType, e.SourceID, e.Description}); err != nil { + return err + } + } + if err := rows.Err(); err != nil { + return err + } + w.Flush() + return w.Error() +} diff --git a/internal/billing/operations_test.go b/internal/billing/operations_test.go new file mode 100644 index 0000000..46bedeb --- /dev/null +++ b/internal/billing/operations_test.go @@ -0,0 +1,36 @@ +package billing + +import ( + "testing" + "time" +) + +func TestOperationalStatusReadiness(t *testing.T) { + now := time.Now().UTC() + recent := now.Add(-time.Minute) + clean := OperationalStatus{ + StripeEnabled: true, ReconciliationStatus: "clean", ReconciliationCompletedAt: &recent, + } + if !clean.Ready(now) { + t.Fatal("clean operational state should be ready") + } + oldWebhook := now.Add(-6 * time.Minute) + stuckWebhook := clean + stuckWebhook.UnprocessedWebhooks = 1 + stuckWebhook.OldestUnprocessedWebhook = &oldWebhook + if stuckWebhook.Ready(now) { + t.Fatal("stuck webhook should fail readiness") + } + freshRefund := now.Add(-time.Minute) + processingRefund := clean + processingRefund.RefundBacklog = 1 + processingRefund.OldestRefund = &freshRefund + if !processingRefund.Ready(now) { + t.Fatal("fresh refund operation should remain ready during its processing window") + } + unmetered := clean + unmetered.UnmeteredSuccesses = 1 + if unmetered.Ready(now) { + t.Fatal("unmetered success must fail readiness") + } +} diff --git a/internal/billing/service.go b/internal/billing/service.go index b87a0bd..30f0e32 100644 --- a/internal/billing/service.go +++ b/internal/billing/service.go @@ -1,6 +1,7 @@ package billing import ( + "bufio" "context" "crypto/rand" "encoding/hex" @@ -8,13 +9,17 @@ import ( "errors" "fmt" "math/big" + "os" + "path/filepath" "strings" + "sync" "time" "aigw/internal/domain" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" + "github.com/stripe/stripe-go/v86" ) const microsPerUnit = int64(1_000_000) @@ -29,8 +34,15 @@ type Service struct { stripeWebhookSecret string stripeSuccessURL string stripeCancelURL string + stripePortalReturnURL string + stripeAutomaticTax bool + stripeProductTaxCode string integrationIdentifier string createStripeCheckout stripeCheckoutCreator + stripeClient *stripe.Client + settlementSpoolPath string + metrics OperationalMetrics + spoolMu sync.Mutex } func New(ctx context.Context, options Options) (*Service, error) { @@ -47,10 +59,15 @@ func New(ctx context.Context, options Options) (*Service, error) { minTopUpMinor: options.MinTopUpMinor, maxTopUpMinor: options.MaxTopUpMinor, stripeEnabled: options.StripeEnabled, stripeWebhookSecret: options.StripeWebhookSecret, stripeSuccessURL: options.StripeSuccessURL, stripeCancelURL: options.StripeCancelURL, + stripePortalReturnURL: options.StripePortalReturnURL, stripeAutomaticTax: options.StripeAutomaticTax, + stripeProductTaxCode: options.StripeProductTaxCode, integrationIdentifier: "aigw_balance_" + randomLetters(8), + settlementSpoolPath: strings.TrimSpace(options.SettlementSpoolPath), + metrics: options.Metrics, } if options.StripeEnabled { - service.createStripeCheckout = newStripeCheckoutCreator(options.StripeAPIKey) + service.stripeClient = stripe.NewClient(options.StripeAPIKey) + service.createStripeCheckout = service.stripeClient.V1CheckoutSessions.Create } return service, nil } @@ -67,10 +84,15 @@ func (s *Service) Currency() string { return s.currency } +func (s *Service) Ping(ctx context.Context) error { return s.db.Ping(ctx) } + func (s *Service) Authorize(ctx context.Context, input Authorization) error { if input.RequestID == "" || input.Principal.TenantID == "" || input.Principal.ProjectID == "" || input.Principal.KeyID == "" { return errors.New("billing authorization identity is incomplete") } + if input.Model.PriceCurrency != "" && input.Model.PriceCurrency != s.currency { + return fmt.Errorf("model price currency %s does not match wallet currency %s", input.Model.PriceCurrency, s.currency) + } reserved, err := reservationCost(input.Model, input.Body, s.defaultMaxOutputTokens) if err != nil { return err @@ -101,7 +123,7 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error { var used, pending int64 if err := tx.QueryRow(ctx, `SELECT COALESCE((SELECT cost_micros FROM usage_monthly_rollups WHERE project_id=$1 AND period_start=$2),0), - COALESCE((SELECT sum(reserved_micros) FROM billing_reservations WHERE project_id=$1 AND status='pending' AND created_at >= $2 AND created_at < $3),0)`, + COALESCE((SELECT sum(reserved_micros) FROM billing_reservations WHERE project_id=$1 AND status IN ('pending','metering_failed') AND created_at >= $2 AND created_at < $3),0)`, input.Principal.ProjectID, period, nextPeriod).Scan(&used, &pending); err != nil { return fmt.Errorf("read monthly spend quota: %w", err) } @@ -115,15 +137,19 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error { if _, err := tx.Exec(ctx, ` INSERT INTO billing_reservations ( request_id, tenant_id, project_id, key_id, public_model, currency, reserved_micros, + price_version_id, input_price_micros_per_million, output_price_micros_per_million, cache_read_price_micros_per_million, cache_write_price_micros_per_million) - VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11)`, + VALUES ($1,$2,$3,$4,$5,$6,$7,NULLIF($8,'')::uuid,$9,$10,$11,$12)`, input.RequestID, input.Principal.TenantID, input.Principal.ProjectID, input.Principal.KeyID, - input.Model.ID, s.currency, reserved, input.Model.InputPriceMicrosPerMillion, + input.Model.ID, s.currency, reserved, input.Model.PriceVersionID, input.Model.InputPriceMicrosPerMillion, input.Model.OutputPriceMicrosPerMillion, input.Model.CacheReadPriceMicrosPerMillion, input.Model.CacheWritePriceMicrosPerMillion); err != nil { return fmt.Errorf("create billing reservation: %w", err) } + if _, err := tx.Exec(ctx, `INSERT INTO billing_settlement_jobs (request_id) VALUES ($1)`, input.RequestID); err != nil { + return fmt.Errorf("create settlement job: %w", err) + } if _, err := tx.Exec(ctx, ` UPDATE tenant_wallets SET reserved_micros = reserved_micros + $2, updated_at = now() WHERE tenant_id = $1`, input.Principal.TenantID, reserved); err != nil { @@ -135,7 +161,27 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error { return nil } -func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error { +func (s *Service) EnqueueSettlement(ctx context.Context, event domain.UsageEvent) error { + payload, err := json.Marshal(event) + if err != nil { + return fmt.Errorf("encode settlement event: %w", err) + } + _, err = s.db.Exec(ctx, ` + INSERT INTO billing_settlement_jobs (request_id, event, status, available_at, updated_at) + VALUES ($1,$2,'pending',now(),now()) + ON CONFLICT (request_id) DO UPDATE SET event=EXCLUDED.event, + status=CASE WHEN billing_settlement_jobs.status='done' THEN 'done' ELSE 'pending' END, + available_at=now(), locked_at=NULL, last_error='', updated_at=now()`, event.RequestID, payload) + if err == nil { + return nil + } + if spoolErr := s.appendSettlementSpool(payload); spoolErr != nil { + return fmt.Errorf("enqueue settlement in PostgreSQL: %v; append durable spool: %w", err, spoolErr) + } + return nil +} + +func (s *Service) settle(ctx context.Context, event domain.UsageEvent) error { tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}) if err != nil { return fmt.Errorf("begin usage settlement: %w", err) @@ -159,7 +205,51 @@ func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error { } actualCost := int64(0) - if event.StatusCode >= 200 && event.StatusCode < 300 { + billableSuccess := event.StatusCode >= 200 && event.StatusCode < 300 && event.Success + if billableSuccess && !event.UsageReported { + // Fail closed: keep the authorization hold in place and make the request + // visible to reconciliation. Releasing it would turn an unmetered success + // into a free request; guessing tokens here could overcharge the customer. + if _, err := tx.Exec(ctx, `UPDATE billing_reservations SET status='metering_failed', settled_at=now() + WHERE request_id=$1`, event.RequestID); err != nil { + return fmt.Errorf("mark unmetered reservation: %w", err) + } + var usageAlreadyRecorded bool + if err := tx.QueryRow(ctx, `SELECT true FROM usage_events WHERE request_id=$1 FOR UPDATE`, event.RequestID).Scan(&usageAlreadyRecorded); err != nil && !errors.Is(err, pgx.ErrNoRows) { + return fmt.Errorf("lock unmetered usage event: %w", err) + } + if _, err := tx.Exec(ctx, `INSERT INTO usage_events ( + request_id,tenant_id,project_id,key_id,public_model,provider_id,upstream_model,protocol,stream, + status_code,success,error_type,attempts,started_at,duration_ms,input_tokens,output_tokens,total_tokens, + cache_creation_input_tokens,cache_read_input_tokens,cost_micros,charged_micros,uncollected_micros, + usage_reported,metering_status) + VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,'usage_not_reported',$12,$13,$14, + $15,$16,$17,$18,$19,0,0,0,false,'missing') + ON CONFLICT (request_id) DO UPDATE SET error_type='usage_not_reported',usage_reported=false,metering_status='missing'`, + event.RequestID, tenantID, projectID, keyID, modelID, event.ProviderID, event.UpstreamModel, + string(event.Protocol), event.Stream, event.StatusCode, event.Success, event.Attempts, event.StartedAt, + event.DurationMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens, + event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens); err != nil { + return fmt.Errorf("persist unmetered usage event: %w", err) + } + if !usageAlreadyRecorded { + period := time.Date(event.StartedAt.UTC().Year(), event.StartedAt.UTC().Month(), 1, 0, 0, 0, 0, time.UTC) + if _, err := tx.Exec(ctx, `INSERT INTO usage_monthly_rollups + (period_start,tenant_id,project_id,request_count,successful_requests,input_tokens,output_tokens,total_tokens) + VALUES ($1,$2,$3,1,1,$4,$5,$6) + ON CONFLICT (project_id,period_start) DO UPDATE SET + request_count=usage_monthly_rollups.request_count+1, + successful_requests=usage_monthly_rollups.successful_requests+1, + input_tokens=usage_monthly_rollups.input_tokens+EXCLUDED.input_tokens, + output_tokens=usage_monthly_rollups.output_tokens+EXCLUDED.output_tokens, + total_tokens=usage_monthly_rollups.total_tokens+EXCLUDED.total_tokens,updated_at=now()`, + period, tenantID, projectID, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens); err != nil { + return fmt.Errorf("roll up unmetered usage event: %w", err) + } + } + return tx.Commit(ctx) + } + if billableSuccess { actualCost, err = usageCost(event.Usage, inputPrice, outputPrice, cacheReadPrice, cacheWritePrice) if err != nil { return err @@ -195,8 +285,9 @@ func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error { return fmt.Errorf("lock existing usage event: %w", err) } if usageAlreadyRecorded { - if _, err := tx.Exec(ctx, `UPDATE usage_events SET cost_micros=$2, charged_micros=$3, uncollected_micros=$4 WHERE request_id=$1`, - event.RequestID, actualCost, charged, uncollected); err != nil { + if _, err := tx.Exec(ctx, `UPDATE usage_events SET cost_micros=$2, charged_micros=$3, uncollected_micros=$4, + usage_reported=$5,metering_status=$6 WHERE request_id=$1`, + event.RequestID, actualCost, charged, uncollected, event.UsageReported, meteringStatus(event)); err != nil { return fmt.Errorf("apply usage charge: %w", err) } } else if _, err := tx.Exec(ctx, ` @@ -204,14 +295,14 @@ func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error { request_id, tenant_id, project_id, key_id, public_model, provider_id, upstream_model, protocol, stream, status_code, success, error_type, attempts, started_at, duration_ms, input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens, - cost_micros, charged_micros, uncollected_micros) - VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23) + cost_micros, charged_micros, uncollected_micros, usage_reported, metering_status) + VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25) ON CONFLICT (request_id) DO NOTHING`, event.RequestID, tenantID, projectID, keyID, modelID, event.ProviderID, event.UpstreamModel, string(event.Protocol), event.Stream, event.StatusCode, event.Success, event.ErrorType, event.Attempts, event.StartedAt, event.DurationMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens, event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens, - actualCost, charged, uncollected); err != nil { + actualCost, charged, uncollected, event.UsageReported, meteringStatus(event)); err != nil { return fmt.Errorf("persist usage event: %w", err) } if charged > 0 { @@ -251,6 +342,199 @@ func (s *Service) Settle(ctx context.Context, event domain.UsageEvent) error { return nil } +// RunSettlementWorker processes jobs using PostgreSQL row locks so any gateway +// instance can resume work left by another instance after a crash. +func (s *Service) RunSettlementWorker(ctx context.Context) { + ticker := time.NewTicker(time.Second) + defer ticker.Stop() + for { + s.drainSettlementSpool(ctx) + s.recoverStaleSettlements(ctx) + for i := 0; i < 32; i++ { + processed, err := s.processSettlementJob(ctx) + if err != nil || !processed { + break + } + } + select { + case <-ctx.Done(): + return + case <-ticker.C: + } + } +} + +func (s *Service) recoverStaleSettlements(ctx context.Context) { + // A process may die after authorization but before it can attach a response + // event. Releasing after an hour prevents permanent holds; such synthetic + // events stay visible in Usage for reconciliation. + rows, err := s.db.Query(ctx, `SELECT request_id,tenant_id::text,project_id::text,key_id::text,public_model,created_at + FROM billing_reservations WHERE status='pending' AND created_at<now()-interval '1 hour' + AND EXISTS(SELECT 1 FROM billing_settlement_jobs j WHERE j.request_id=billing_reservations.request_id AND j.status='awaiting_event') LIMIT 100`) + if err != nil { + return + } + defer rows.Close() + for rows.Next() { + var event domain.UsageEvent + if rows.Scan(&event.RequestID, &event.TenantID, &event.ProjectID, &event.KeyID, &event.PublicModel, &event.StartedAt) != nil { + continue + } + event.StatusCode = 500 + event.Success = false + event.ErrorType = "gateway_interrupted_before_settlement" + event.DurationMS = time.Since(event.StartedAt).Milliseconds() + _ = s.EnqueueSettlement(ctx, event) + } + _, _ = s.db.Exec(ctx, `UPDATE billing_settlement_jobs SET status='retry',locked_at=NULL,available_at=now(),updated_at=now(),last_error='recovered stale processing lease' + WHERE status='processing' AND locked_at<now()-interval '5 minutes'`) +} + +func (s *Service) processSettlementJob(ctx context.Context) (bool, error) { + tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}) + if err != nil { + return false, err + } + defer tx.Rollback(ctx) + var requestID string + var payload []byte + err = tx.QueryRow(ctx, ` + WITH selected AS ( + SELECT request_id FROM billing_settlement_jobs + WHERE status IN ('pending','retry') AND available_at <= now() + ORDER BY available_at, created_at FOR UPDATE SKIP LOCKED LIMIT 1 + ) + UPDATE billing_settlement_jobs j SET status='processing', attempts=attempts+1, + locked_at=now(), updated_at=now() + FROM selected WHERE j.request_id=selected.request_id + RETURNING j.request_id, j.event`, + ).Scan(&requestID, &payload) + if errors.Is(err, pgx.ErrNoRows) { + return false, tx.Commit(ctx) + } + if err != nil { + return false, err + } + if err := tx.Commit(ctx); err != nil { + return false, err + } + var event domain.UsageEvent + if err := json.Unmarshal(payload, &event); err != nil { + s.retrySettlement(ctx, requestID, fmt.Errorf("decode settlement event: %w", err)) + return true, err + } + jobCtx, cancel := context.WithTimeout(ctx, 15*time.Second) + err = s.settle(jobCtx, event) + cancel() + if err != nil { + s.retrySettlement(ctx, requestID, err) + return true, err + } + _, err = s.db.Exec(ctx, `UPDATE billing_settlement_jobs SET status='done', locked_at=NULL, + last_error='', updated_at=now(), completed_at=now() WHERE request_id=$1`, requestID) + return true, err +} + +func (s *Service) retrySettlement(ctx context.Context, requestID string, cause error) { + message := cause.Error() + if len(message) > 1000 { + message = message[:1000] + } + _, _ = s.db.Exec(ctx, `UPDATE billing_settlement_jobs SET status='retry', locked_at=NULL, + last_error=$2, available_at=now() + make_interval(secs => LEAST(300, power(2, LEAST(attempts, 8))::int)), + updated_at=now() WHERE request_id=$1`, requestID, message) +} + +func (s *Service) appendSettlementSpool(payload []byte) error { + if s.settlementSpoolPath == "" { + return errors.New("settlement spool path is not configured") + } + s.spoolMu.Lock() + defer s.spoolMu.Unlock() + if err := os.MkdirAll(filepath.Dir(s.settlementSpoolPath), 0o700); err != nil { + return err + } + file, err := os.OpenFile(s.settlementSpoolPath, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o600) + if err != nil { + return err + } + defer file.Close() + if _, err := file.Write(append(payload, '\n')); err != nil { + return err + } + return file.Sync() +} + +func (s *Service) drainSettlementSpool(ctx context.Context) { + if s.settlementSpoolPath == "" { + return + } + s.spoolMu.Lock() + defer s.spoolMu.Unlock() + file, err := os.Open(s.settlementSpoolPath) + if errors.Is(err, os.ErrNotExist) { + return + } + if err != nil { + return + } + var pending [][]byte + scanner := bufio.NewScanner(file) + scanner.Buffer(make([]byte, 64*1024), 2<<20) + for scanner.Scan() { + line := append([]byte(nil), scanner.Bytes()...) + var event domain.UsageEvent + if json.Unmarshal(line, &event) != nil || event.RequestID == "" { + pending = append(pending, line) + continue + } + if _, err := s.db.Exec(ctx, `UPDATE billing_settlement_jobs SET event=$2, status=CASE WHEN status='done' THEN 'done' ELSE 'pending' END, + available_at=now(), locked_at=NULL, updated_at=now() WHERE request_id=$1`, event.RequestID, line); err != nil { + pending = append(pending, line) + } + } + _ = file.Close() + if scanner.Err() != nil { + return + } + temporary := s.settlementSpoolPath + ".tmp" + out, err := os.OpenFile(temporary, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0o600) + if err != nil { + return + } + for _, line := range pending { + _, _ = out.Write(append(line, '\n')) + } + _ = out.Sync() + _ = out.Close() + _ = os.Rename(temporary, s.settlementSpoolPath) +} + +func (s *Service) SettlementQueueStatus(ctx context.Context) (SettlementQueueStatus, error) { + var result SettlementQueueStatus + err := s.db.QueryRow(ctx, `SELECT + count(*) FILTER (WHERE status='awaiting_event'), count(*) FILTER (WHERE status='pending'), + count(*) FILTER (WHERE status='processing'), count(*) FILTER (WHERE status='retry'), + min(created_at) FILTER (WHERE status IN ('awaiting_event','pending','processing','retry')) + FROM billing_settlement_jobs`).Scan(&result.AwaitingEvent, &result.Pending, &result.Processing, &result.Retrying, &result.OldestPending) + if err != nil { + return result, err + } + if s.settlementSpoolPath != "" { + s.spoolMu.Lock() + file, openErr := os.Open(s.settlementSpoolPath) + if openErr == nil { + scanner := bufio.NewScanner(file) + for scanner.Scan() { + result.SpoolRecords++ + } + _ = file.Close() + } + s.spoolMu.Unlock() + } + return result, nil +} + func boolToInt(value bool) int { if value { return 1 @@ -258,8 +542,21 @@ func boolToInt(value bool) int { return 0 } +func meteringStatus(event domain.UsageEvent) string { + if event.StatusCode < 200 || event.StatusCode >= 300 || !event.Success { + return "upstream_failed" + } + if event.UsageReported { + return "reported" + } + return "missing" +} + func reservationCost(model domain.Model, body []byte, defaultMaxOutput int64) (int64, error) { maxOutput := defaultMaxOutput + if model.MaxOutputTokens > 0 && (maxOutput == 0 || model.MaxOutputTokens < maxOutput) { + maxOutput = model.MaxOutputTokens + } var limits struct { MaxTokens int64 `json:"max_tokens"` MaxCompletionTokens int64 `json:"max_completion_tokens"` diff --git a/internal/billing/service_test.go b/internal/billing/service_test.go index f675da7..6a1bb49 100644 --- a/internal/billing/service_test.go +++ b/internal/billing/service_test.go @@ -2,6 +2,7 @@ package billing import ( "context" + "crypto/sha256" "encoding/json" "fmt" "net/http" @@ -97,6 +98,26 @@ func TestCheckoutReturnURLPreservesCallbackAndSessionPlaceholder(t *testing.T) { } } +func TestSettlementSpoolWritesOneDurableRecord(t *testing.T) { + path := t.TempDir() + "/settlements.jsonl" + service := &Service{settlementSpoolPath: path} + event := domain.UsageEvent{RequestID: "req_spool_test", TenantID: "tenant", StartedAt: time.Now().UTC()} + payload, err := json.Marshal(event) + if err != nil { + t.Fatal(err) + } + if err := service.appendSettlementSpool(payload); err != nil { + t.Fatal(err) + } + contents, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + if strings.Count(string(contents), "req_spool_test") != 1 || !strings.HasSuffix(string(contents), "\n") { + t.Fatalf("unexpected spool contents %q", contents) + } +} + func TestWebhookRejectsInvalidSignatureBeforeProcessing(t *testing.T) { service := &Service{stripeWebhookSecret: "whsec_test"} request := httptest.NewRequest(http.MethodPost, "/billing/stripe/webhook", strings.NewReader(`{"id":"evt_fake"}`)) @@ -133,6 +154,7 @@ func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) { t.Fatal(err) } eventID := fmt.Sprintf("evt_aigw_%d", time.Now().UnixNano()) + followupEventID := eventID + "_async" t.Cleanup(func() { cleanupCtx := context.Background() for _, statement := range []struct { @@ -140,6 +162,7 @@ func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) { arg string }{ {`DELETE FROM stripe_webhook_events WHERE event_id=$1`, eventID}, + {`DELETE FROM stripe_webhook_events WHERE event_id=$1`, followupEventID}, {`DELETE FROM billing_ledger WHERE tenant_id=$1`, tenantID}, {`DELETE FROM tenant_wallets WHERE tenant_id=$1`, tenantID}, {`DELETE FROM topup_orders WHERE tenant_id=$1`, tenantID}, @@ -177,6 +200,28 @@ func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) { t.Fatalf("delivery %d status = %d, body = %s", delivery+1, response.Code, response.Body.String()) } } + if _, err := service.db.Exec(ctx, `UPDATE topup_orders SET status='refunded',refunded_micros=amount_micros WHERE id=$1`, orderID); err != nil { + t.Fatal(err) + } + followupPayload, err := json.Marshal(map[string]any{ + "id": followupEventID, "object": "event", "api_version": stripe.APIVersion, + "type": string(stripe.EventTypeCheckoutSessionAsyncPaymentSucceeded), + "data": map[string]any{"object": map[string]any{ + "id": sessionID, "object": "checkout.session", "client_reference_id": orderID, + "amount_total": amountMinor, "currency": "usd", "payment_status": "paid", + }}, + }) + if err != nil { + t.Fatal(err) + } + followupSigned := webhook.GenerateTestSignedPayload(&webhook.UnsignedPayload{Payload: followupPayload, Secret: service.stripeWebhookSecret}) + followupRequest := httptest.NewRequest(http.MethodPost, "/billing/stripe/webhook", strings.NewReader(string(followupPayload))) + followupRequest.Header.Set("Stripe-Signature", followupSigned.Header) + followupResponse := httptest.NewRecorder() + service.WebhookHandler().ServeHTTP(followupResponse, followupRequest) + if followupResponse.Code != http.StatusOK { + t.Fatalf("follow-up status = %d, body = %s", followupResponse.Code, followupResponse.Body.String()) + } var balance int64 if err := service.db.QueryRow(ctx, `SELECT balance_micros FROM tenant_wallets WHERE tenant_id=$1`, tenantID).Scan(&balance); err != nil { @@ -190,13 +235,136 @@ func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) { if err := service.db.QueryRow(ctx, `SELECT count(*) FROM billing_ledger WHERE source_type='stripe_checkout' AND source_id=$1`, sessionID).Scan(&ledgerCount); err != nil { t.Fatal(err) } - if err := service.db.QueryRow(ctx, `SELECT count(*) FROM stripe_webhook_events WHERE event_id=$1`, eventID).Scan(&webhookCount); err != nil { + if err := service.db.QueryRow(ctx, `SELECT count(*) FROM stripe_webhook_events WHERE event_id IN ($1,$2)`, eventID, followupEventID).Scan(&webhookCount); err != nil { t.Fatal(err) } if err := service.db.QueryRow(ctx, `SELECT status FROM topup_orders WHERE id=$1`, orderID).Scan(&orderStatus); err != nil { t.Fatal(err) } - if ledgerCount != 1 || webhookCount != 1 || orderStatus != "paid" { - t.Fatalf("ledger=%d webhook=%d order=%s, want 1/1/paid", ledgerCount, webhookCount, orderStatus) + if ledgerCount != 1 || webhookCount != 2 || orderStatus != "refunded" { + t.Fatalf("ledger=%d webhook=%d order=%s, want 1/2/refunded", ledgerCount, webhookCount, orderStatus) + } +} + +func TestSettlementWorkerPersistsUsageAndReleasesReservationPostgres(t *testing.T) { + databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL") + if databaseURL == "" { + t.Skip("AIGW_TEST_DATABASE_URL is not set") + } + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + if err := controlplane.MigrateDatabase(ctx, databaseURL); err != nil { + t.Fatal(err) + } + service, err := New(ctx, Options{DatabaseURL: databaseURL, Currency: "usd", SettlementSpoolPath: t.TempDir() + "/settlements.jsonl"}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(service.Close) + slug := fmt.Sprintf("settle-%d", time.Now().UnixNano()) + var tenantID, projectID, keyID string + if err := service.db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Settlement integration') RETURNING id::text`, slug).Scan(&tenantID); err != nil { + t.Fatal(err) + } + if err := service.db.QueryRow(ctx, `INSERT INTO projects (tenant_id,slug,name) VALUES ($1,'default','Settlement') RETURNING id::text`, tenantID).Scan(&projectID); err != nil { + t.Fatal(err) + } + keyHash := sha256.Sum256([]byte(slug)) + if err := service.db.QueryRow(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_hash) VALUES ($1,$2,'integration','sk-test',$3) RETURNING id::text`, tenantID, projectID, keyHash[:]).Scan(&keyID); err != nil { + t.Fatal(err) + } + requestID := fmt.Sprintf("req_settle_%d", time.Now().UnixNano()) + missingRequestID := requestID + "_missing_usage" + t.Cleanup(func() { + for _, statement := range []struct { + query string + arg string + }{ + {`DELETE FROM billing_ledger WHERE tenant_id=$1`, tenantID}, + {`DELETE FROM usage_monthly_rollups WHERE tenant_id=$1`, tenantID}, + {`DELETE FROM usage_events WHERE tenant_id=$1`, tenantID}, + {`DELETE FROM billing_settlement_jobs WHERE request_id=$1`, requestID}, + {`DELETE FROM billing_settlement_jobs WHERE request_id=$1`, missingRequestID}, + {`DELETE FROM billing_reservations WHERE request_id=$1`, requestID}, + {`DELETE FROM billing_reservations WHERE request_id=$1`, missingRequestID}, + {`DELETE FROM tenant_wallets WHERE tenant_id=$1`, tenantID}, + {`DELETE FROM api_keys WHERE id=$1`, keyID}, + {`DELETE FROM projects WHERE id=$1`, projectID}, + {`DELETE FROM tenants WHERE id=$1`, tenantID}, + } { + if _, cleanupErr := service.db.Exec(context.Background(), statement.query, statement.arg); cleanupErr != nil { + t.Errorf("cleanup settlement integration data: %v", cleanupErr) + } + } + }) + if _, err := service.db.Exec(ctx, `INSERT INTO tenant_wallets (tenant_id,currency,balance_micros,reserved_micros) VALUES ($1,'usd',1000,20)`, tenantID); err != nil { + t.Fatal(err) + } + if _, err := service.db.Exec(ctx, `INSERT INTO billing_reservations (request_id,tenant_id,project_id,key_id,public_model,currency,reserved_micros,input_price_micros_per_million,output_price_micros_per_million,cache_read_price_micros_per_million,cache_write_price_micros_per_million) VALUES ($1,$2,$3,$4,'demo/model','usd',20,1000000,1000000,0,0)`, requestID, tenantID, projectID, keyID); err != nil { + t.Fatal(err) + } + if _, err := service.db.Exec(ctx, `INSERT INTO billing_settlement_jobs (request_id) VALUES ($1)`, requestID); err != nil { + t.Fatal(err) + } + if err := service.EnqueueSettlement(ctx, domain.UsageEvent{RequestID: requestID, TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "demo/model", Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: time.Now().UTC(), Usage: domain.Usage{InputTokens: 10}}); err != nil { + t.Fatal(err) + } + processed, err := service.processSettlementJob(ctx) + if err != nil || !processed { + t.Fatalf("process settlement: processed=%v err=%v", processed, err) + } + var reservationStatus, jobStatus string + var balance, reserved, charged, usageCount int64 + if err := service.db.QueryRow(ctx, `SELECT status,charged_micros FROM billing_reservations WHERE request_id=$1`, requestID).Scan(&reservationStatus, &charged); err != nil { + t.Fatal(err) + } + if err := service.db.QueryRow(ctx, `SELECT status FROM billing_settlement_jobs WHERE request_id=$1`, requestID).Scan(&jobStatus); err != nil { + t.Fatal(err) + } + if err := service.db.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1`, tenantID).Scan(&balance, &reserved); err != nil { + t.Fatal(err) + } + if err := service.db.QueryRow(ctx, `SELECT count(*) FROM usage_events WHERE request_id=$1`, requestID).Scan(&usageCount); err != nil { + t.Fatal(err) + } + if reservationStatus != "settled" || jobStatus != "done" || balance != 990 || reserved != 0 || charged != 10 || usageCount != 1 { + t.Fatalf("reservation=%s job=%s balance=%d reserved=%d charged=%d usage=%d", reservationStatus, jobStatus, balance, reserved, charged, usageCount) + } + + if _, err := service.db.Exec(ctx, `UPDATE tenant_wallets SET reserved_micros=reserved_micros+30 WHERE tenant_id=$1`, tenantID); err != nil { + t.Fatal(err) + } + if _, err := service.db.Exec(ctx, `INSERT INTO billing_reservations + (request_id,tenant_id,project_id,key_id,public_model,currency,reserved_micros,input_price_micros_per_million,output_price_micros_per_million,cache_read_price_micros_per_million,cache_write_price_micros_per_million) + VALUES ($1,$2,$3,$4,'demo/model','usd',30,1000000,1000000,0,0)`, missingRequestID, tenantID, projectID, keyID); err != nil { + t.Fatal(err) + } + if _, err := service.db.Exec(ctx, `INSERT INTO billing_settlement_jobs (request_id) VALUES ($1)`, missingRequestID); err != nil { + t.Fatal(err) + } + if err := service.EnqueueSettlement(ctx, domain.UsageEvent{RequestID: missingRequestID, TenantID: tenantID, ProjectID: projectID, + KeyID: keyID, PublicModel: "demo/model", Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, + UsageReported: false, StartedAt: time.Now().UTC()}); err != nil { + t.Fatal(err) + } + processed, err = service.processSettlementJob(ctx) + if err != nil || !processed { + t.Fatalf("process missing-usage settlement: processed=%v err=%v", processed, err) + } + var missingReservationStatus, missingJobStatus, meteringStatus string + if err := service.db.QueryRow(ctx, `SELECT status FROM billing_reservations WHERE request_id=$1`, missingRequestID).Scan(&missingReservationStatus); err != nil { + t.Fatal(err) + } + if err := service.db.QueryRow(ctx, `SELECT status FROM billing_settlement_jobs WHERE request_id=$1`, missingRequestID).Scan(&missingJobStatus); err != nil { + t.Fatal(err) + } + if err := service.db.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1`, tenantID).Scan(&balance, &reserved); err != nil { + t.Fatal(err) + } + if err := service.db.QueryRow(ctx, `SELECT metering_status FROM usage_events WHERE request_id=$1`, missingRequestID).Scan(&meteringStatus); err != nil { + t.Fatal(err) + } + if missingReservationStatus != "metering_failed" || missingJobStatus != "done" || balance != 990 || reserved != 30 || meteringStatus != "missing" { + t.Fatalf("missing usage reservation=%s job=%s balance=%d reserved=%d metering=%s", missingReservationStatus, missingJobStatus, balance, reserved, meteringStatus) } } diff --git a/internal/billing/stripe.go b/internal/billing/stripe.go index b887035..59ee9bd 100644 --- a/internal/billing/stripe.go +++ b/internal/billing/stripe.go @@ -39,6 +39,13 @@ func (s *Service) CreateCheckout(ctx context.Context, input CheckoutInput) (Chec "aigw_topup_order_id": orderID, "aigw_tenant_id": strings.TrimSpace(input.TenantID), }, + InvoiceCreation: &stripe.CheckoutSessionCreateInvoiceCreationParams{ + Enabled: stripe.Bool(true), + InvoiceData: &stripe.CheckoutSessionCreateInvoiceCreationInvoiceDataParams{ + Description: stripe.String("AIGW prepaid API usage credit"), + Metadata: map[string]string{"aigw_topup_order_id": orderID, "aigw_tenant_id": strings.TrimSpace(input.TenantID)}, + }, + }, LineItems: []*stripe.CheckoutSessionCreateLineItemParams{{ Quantity: stripe.Int64(1), PriceData: &stripe.CheckoutSessionCreateLineItemPriceDataParams{ @@ -51,10 +58,27 @@ func (s *Service) CreateCheckout(ctx context.Context, input CheckoutInput) (Chec }, }}, } + var customerID string + _ = s.db.QueryRow(ctx, `SELECT stripe_customer_id FROM stripe_customers WHERE tenant_id=$1`, input.TenantID).Scan(&customerID) + if customerID != "" { + params.Customer = stripe.String(customerID) + } else { + params.CustomerCreation = stripe.String(string(stripe.CheckoutSessionCustomerCreationAlways)) + if strings.TrimSpace(input.CustomerEmail) != "" { + params.CustomerEmail = stripe.String(strings.TrimSpace(input.CustomerEmail)) + } + } + if s.stripeAutomaticTax { + params.AutomaticTax = &stripe.CheckoutSessionCreateAutomaticTaxParams{Enabled: stripe.Bool(true)} + params.TaxIDCollection = &stripe.CheckoutSessionCreateTaxIDCollectionParams{Enabled: stripe.Bool(true)} + params.LineItems[0].PriceData.ProductData.TaxCode = stripe.String(s.stripeProductTaxCode) + } params.SetIdempotencyKey("aigw_topup_" + orderID) session, err := s.createStripeCheckout(ctx, params) if err != nil { - _, _ = s.db.Exec(ctx, `UPDATE topup_orders SET status = 'failed' WHERE id = $1 AND status = 'pending'`, orderID) + if _, updateErr := s.db.Exec(ctx, `UPDATE topup_orders SET status = 'failed' WHERE id = $1 AND status = 'pending'`, orderID); updateErr != nil { + return CheckoutResult{}, errors.Join(fmt.Errorf("create Stripe Checkout Session: %w", err), fmt.Errorf("mark top-up order failed: %w", updateErr)) + } return CheckoutResult{}, fmt.Errorf("create Stripe Checkout Session: %w", err) } if session.ID == "" || session.URL == "" { @@ -103,7 +127,12 @@ func (s *Service) WebhookHandler() http.Handler { http.Error(w, "invalid webhook signature", http.StatusBadRequest) return } + if err := s.recordWebhookAttempt(r.Context(), event); err != nil { + http.Error(w, "webhook persistence failed", http.StatusInternalServerError) + return + } if err := s.processStripeEvent(r.Context(), event); err != nil { + s.recordWebhookFailure(r.Context(), event.ID, err) if errors.Is(err, ErrInvalidAmount) || isNotFound(err) { http.Error(w, "invalid checkout event", http.StatusBadRequest) return @@ -117,15 +146,29 @@ func (s *Service) WebhookHandler() http.Handler { } func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) error { - typeName := string(event.Type) switch event.Type { case stripe.EventTypeCheckoutSessionCompleted, stripe.EventTypeCheckoutSessionAsyncPaymentSucceeded, stripe.EventTypeCheckoutSessionAsyncPaymentFailed, stripe.EventTypeCheckoutSessionExpired: + return s.processCheckoutEvent(ctx, event) + case stripe.EventTypeChargeSucceeded, stripe.EventTypeChargeUpdated, stripe.EventTypeChargeRefunded, + stripe.EventTypeRefundCreated, stripe.EventTypeRefundUpdated, stripe.EventTypeRefundFailed, + stripe.EventTypeChargeDisputeCreated, stripe.EventTypeChargeDisputeUpdated, + stripe.EventTypeChargeDisputeClosed, stripe.EventTypeChargeDisputeFundsWithdrawn, + stripe.EventTypeChargeDisputeFundsReinstated, + stripe.EventTypeInvoiceCreated, stripe.EventTypeInvoiceFinalized, + stripe.EventTypeInvoicePaid, stripe.EventTypeInvoicePaymentFailed: + return s.processOperationalStripeEvent(ctx, event) default: - return nil + _, err := s.db.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at=now(),processing_error='ignored event type' + WHERE event_id=$1 AND processed_at IS NULL`, event.ID) + return err } +} + +func (s *Service) processCheckoutEvent(ctx context.Context, event stripe.Event) error { + typeName := string(event.Type) if event.Data == nil { return ErrInvalidAmount } @@ -141,6 +184,9 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er return err } defer tx.Rollback(ctx) + if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1,0))`, event.ID); err != nil { + return err + } tag, err := tx.Exec(ctx, ` INSERT INTO stripe_webhook_events (event_id, event_type) VALUES ($1,$2) ON CONFLICT (event_id) DO NOTHING`, event.ID, typeName) @@ -148,7 +194,13 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er return fmt.Errorf("record Stripe event: %w", err) } if tag.RowsAffected() == 0 { - return tx.Commit(ctx) + var processed bool + if err := tx.QueryRow(ctx, `SELECT processed_at IS NOT NULL FROM stripe_webhook_events WHERE event_id=$1`, event.ID).Scan(&processed); err != nil { + return err + } + if processed { + return tx.Commit(ctx) + } } var tenantID, currency, status string @@ -164,6 +216,32 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er if (storedSessionID != nil && *storedSessionID != session.ID) || amountMinor != session.AmountTotal || currency != string(session.Currency) { return ErrInvalidAmount } + customerID, paymentIntentID, invoiceID := "", "", "" + if session.Customer != nil { + customerID = session.Customer.ID + } + if session.PaymentIntent != nil { + paymentIntentID = session.PaymentIntent.ID + } + if session.Invoice != nil { + invoiceID = session.Invoice.ID + } + if customerID != "" { + email := session.CustomerEmail + if session.CustomerDetails != nil && session.CustomerDetails.Email != "" { + email = session.CustomerDetails.Email + } + if _, err := tx.Exec(ctx, `INSERT INTO stripe_customers (tenant_id,stripe_customer_id,email) VALUES ($1,$2,$3) + ON CONFLICT (tenant_id) DO UPDATE SET stripe_customer_id=EXCLUDED.stripe_customer_id, + email=CASE WHEN EXCLUDED.email='' THEN stripe_customers.email ELSE EXCLUDED.email END,updated_at=now()`, tenantID, customerID, email); err != nil { + return err + } + } + if _, err := tx.Exec(ctx, `UPDATE topup_orders SET stripe_customer_id=COALESCE(NULLIF($2,''),stripe_customer_id), + stripe_payment_intent_id=COALESCE(NULLIF($3,''),stripe_payment_intent_id), + stripe_invoice_id=COALESCE(NULLIF($4,''),stripe_invoice_id) WHERE id=$1`, session.ClientReferenceID, customerID, paymentIntentID, invoiceID); err != nil { + return err + } if event.Type == stripe.EventTypeCheckoutSessionAsyncPaymentFailed || event.Type == stripe.EventTypeCheckoutSessionExpired { orderStatus := "failed" if event.Type == stripe.EventTypeCheckoutSessionExpired { @@ -172,20 +250,24 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er if _, err := tx.Exec(ctx, `UPDATE topup_orders SET status = $2, stripe_session_id = COALESCE(stripe_session_id, $3) WHERE id = $1 AND status = 'pending'`, session.ClientReferenceID, orderStatus, session.ID); err != nil { return err } - _, err = tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at = now() WHERE event_id = $1`, event.ID) + _, err = tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at=now(),processing_error='' WHERE event_id=$1`, event.ID) if err != nil { return err } return tx.Commit(ctx) } if session.PaymentStatus != stripe.CheckoutSessionPaymentStatusPaid { - _, err = tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at = now() WHERE event_id = $1`, event.ID) + _, err = tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at=now(),processing_error='' WHERE event_id=$1`, event.ID) if err != nil { return err } return tx.Commit(ctx) } - if status != "paid" { + var alreadyCredited bool + if err := tx.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM billing_ledger WHERE source_type='stripe_checkout' AND source_id=$1)`, session.ID).Scan(&alreadyCredited); err != nil { + return err + } + if !alreadyCredited { if _, err := tx.Exec(ctx, ` INSERT INTO tenant_wallets (tenant_id, currency) VALUES ($1,$2) ON CONFLICT (tenant_id) DO NOTHING`, tenantID, currency); err != nil { @@ -209,14 +291,117 @@ func (s *Service) processStripeEvent(ctx context.Context, event stripe.Event) er ON CONFLICT (source_type, source_id) DO NOTHING`, tenantID, currency, amountMicros, newBalance, session.ID); err != nil { return err } - if _, err := tx.Exec(ctx, ` - UPDATE topup_orders SET status = 'paid', stripe_session_id = COALESCE(stripe_session_id, $2), paid_at = now() - WHERE id = $1`, session.ClientReferenceID, session.ID); err != nil { + } + if _, err := tx.Exec(ctx, ` + UPDATE topup_orders SET status=CASE WHEN status IN ('pending','failed','expired') THEN 'paid' ELSE status END, + stripe_session_id=COALESCE(stripe_session_id,$2), paid_at=COALESCE(paid_at,now()) + WHERE id=$1`, session.ClientReferenceID, session.ID); err != nil { + return err + } + if _, err := tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at=now(),processing_error='' WHERE event_id=$1`, event.ID); err != nil { + return err + } + return tx.Commit(ctx) +} + +func (s *Service) processOperationalStripeEvent(ctx context.Context, event stripe.Event) error { + if event.ID == "" || event.Data == nil { + return ErrInvalidAmount + } + tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}) + if err != nil { + return err + } + defer tx.Rollback(ctx) + if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock(hashtextextended($1,0))`, event.ID); err != nil { + return err + } + tag, err := tx.Exec(ctx, `INSERT INTO stripe_webhook_events (event_id,event_type) VALUES ($1,$2) ON CONFLICT DO NOTHING`, event.ID, string(event.Type)) + if err != nil { + return err + } + if tag.RowsAffected() == 0 { + var processed bool + if err := tx.QueryRow(ctx, `SELECT processed_at IS NOT NULL FROM stripe_webhook_events WHERE event_id=$1`, event.ID).Scan(&processed); err != nil { return err } + if processed { + return tx.Commit(ctx) + } } - if _, err := tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at = now() WHERE event_id = $1`, event.ID); err != nil { + switch event.Type { + case stripe.EventTypeChargeSucceeded, stripe.EventTypeChargeUpdated, stripe.EventTypeChargeRefunded: + var charge stripe.Charge + if json.Unmarshal(event.Data.Raw, &charge) != nil || charge.ID == "" { + return ErrInvalidAmount + } + paymentIntentID := "" + if charge.PaymentIntent != nil { + paymentIntentID = charge.PaymentIntent.ID + } + if paymentIntentID != "" { + if _, err := tx.Exec(ctx, `UPDATE topup_orders SET stripe_charge_id=$2,receipt_url=COALESCE(NULLIF($3,''),receipt_url) WHERE stripe_payment_intent_id=$1`, paymentIntentID, charge.ID, charge.ReceiptURL); err != nil { + return err + } + } + if event.Type == stripe.EventTypeChargeRefunded && charge.Refunds != nil { + for _, refund := range charge.Refunds.Data { + if err := s.applyRefundTx(ctx, tx, refund); err != nil { + return err + } + } + } + case stripe.EventTypeRefundCreated, stripe.EventTypeRefundUpdated, stripe.EventTypeRefundFailed: + var refund stripe.Refund + if json.Unmarshal(event.Data.Raw, &refund) != nil || refund.ID == "" { + return ErrInvalidAmount + } + if err := s.applyRefundTx(ctx, tx, &refund); err != nil { + return err + } + case stripe.EventTypeChargeDisputeCreated, stripe.EventTypeChargeDisputeUpdated, + stripe.EventTypeChargeDisputeClosed, stripe.EventTypeChargeDisputeFundsWithdrawn, stripe.EventTypeChargeDisputeFundsReinstated: + var dispute stripe.Dispute + if json.Unmarshal(event.Data.Raw, &dispute) != nil || dispute.ID == "" { + return ErrInvalidAmount + } + if err := s.applyDisputeTx(ctx, tx, &dispute, event.Type); err != nil { + return err + } + case stripe.EventTypeInvoiceCreated, stripe.EventTypeInvoiceFinalized, stripe.EventTypeInvoicePaid, stripe.EventTypeInvoicePaymentFailed: + var invoice stripe.Invoice + if json.Unmarshal(event.Data.Raw, &invoice) != nil || invoice.ID == "" { + return ErrInvalidAmount + } + if err := s.applyInvoiceTx(ctx, tx, &invoice); err != nil { + return err + } + } + if _, err := tx.Exec(ctx, `UPDATE stripe_webhook_events SET processed_at=now(),processing_error='' WHERE event_id=$1`, event.ID); err != nil { return err } return tx.Commit(ctx) } + +func (s *Service) recordWebhookAttempt(ctx context.Context, event stripe.Event) error { + if event.ID == "" { + return ErrInvalidAmount + } + _, err := s.db.Exec(ctx, `INSERT INTO stripe_webhook_events + (event_id,event_type,attempts,last_attempt_at) VALUES ($1,$2,1,now()) + ON CONFLICT (event_id) DO UPDATE SET attempts=stripe_webhook_events.attempts+1,last_attempt_at=now()`, + event.ID, string(event.Type)) + return err +} + +func (s *Service) recordWebhookFailure(ctx context.Context, eventID string, cause error) { + message := "webhook processing failed" + if cause != nil { + message = cause.Error() + } + if len(message) > 1000 { + message = message[:1000] + } + _, _ = s.db.Exec(ctx, `UPDATE stripe_webhook_events SET processing_error=$2,last_attempt_at=now() + WHERE event_id=$1 AND processed_at IS NULL`, eventID, message) +} diff --git a/internal/billing/types.go b/internal/billing/types.go index 3d0461e..fe5df1e 100644 --- a/internal/billing/types.go +++ b/internal/billing/types.go @@ -14,11 +14,13 @@ var ( ErrInvalidAmount = errors.New("invalid amount") ErrQuotaExceeded = errors.New("monthly spend quota exceeded") ErrTopUpOrderNotFound = errors.New("top-up order not found") + ErrUsageNotReported = errors.New("billable successful response did not report usage") + ErrCannotResolveTopUp = errors.New("top-up order cannot be resolved as missing") ) type Meter interface { Authorize(context.Context, Authorization) error - Settle(context.Context, domain.UsageEvent) error + EnqueueSettlement(context.Context, domain.UsageEvent) error } type Authorization struct { @@ -40,6 +42,54 @@ type Options struct { StripeWebhookSecret string StripeSuccessURL string StripeCancelURL string + StripePortalReturnURL string + StripeAutomaticTax bool + StripeProductTaxCode string + SettlementSpoolPath string + Metrics OperationalMetrics +} + +type OperationalMetrics interface { + SetStripeOperations(OperationalStatus) +} + +type OperationalStatus struct { + StripeEnabled bool `json:"stripe_enabled"` + ReconciliationStatus string `json:"reconciliation_status"` + ReconciliationMismatches int64 `json:"reconciliation_mismatches"` + ReconciliationCompletedAt *time.Time `json:"reconciliation_completed_at,omitempty"` + ReconciliationError string `json:"reconciliation_error,omitempty"` + UnprocessedWebhooks int64 `json:"unprocessed_webhooks"` + OldestUnprocessedWebhook *time.Time `json:"oldest_unprocessed_webhook,omitempty"` + RefundBacklog int64 `json:"refund_backlog"` + OldestRefund *time.Time `json:"oldest_refund,omitempty"` + UncollectedMicros int64 `json:"uncollected_micros"` + UnmeteredSuccesses int64 `json:"unmetered_successes"` +} + +func (s OperationalStatus) Ready(now time.Time) bool { + if !s.StripeEnabled { + return s.UnmeteredSuccesses == 0 + } + if s.ReconciliationCompletedAt == nil || now.Sub(*s.ReconciliationCompletedAt) > 2*time.Hour { + return false + } + webhooksStuck := s.UnprocessedWebhooks > 0 && + (s.OldestUnprocessedWebhook == nil || now.Sub(*s.OldestUnprocessedWebhook) > 5*time.Minute) + refundsStuck := s.RefundBacklog > 0 && + (s.OldestRefund == nil || now.Sub(*s.OldestRefund) > 15*time.Minute) + return s.ReconciliationStatus == "clean" && s.ReconciliationMismatches == 0 && + s.ReconciliationError == "" && !webhooksStuck && !refundsStuck && + s.UncollectedMicros == 0 && s.UnmeteredSuccesses == 0 +} + +type SettlementQueueStatus struct { + AwaitingEvent int64 `json:"awaiting_event"` + Pending int64 `json:"pending"` + Processing int64 `json:"processing"` + Retrying int64 `json:"retrying"` + OldestPending *time.Time `json:"oldest_pending,omitempty"` + SpoolRecords int `json:"spool_records"` } type Account struct { @@ -73,8 +123,9 @@ type AdjustmentInput struct { } type CheckoutInput struct { - TenantID string `json:"tenant_id"` - AmountMinor int64 `json:"amount_minor"` + TenantID string `json:"tenant_id"` + AmountMinor int64 `json:"amount_minor"` + CustomerEmail string `json:"-"` } type CheckoutResult struct { @@ -84,14 +135,99 @@ type CheckoutResult struct { } type TopUpOrder struct { - ID string `json:"id"` - TenantID string `json:"tenant_id"` - AmountMinor int64 `json:"amount_minor"` - AmountMicros int64 `json:"amount_micros"` - Currency string `json:"currency"` - Status string `json:"status"` - StripeSessionID string `json:"stripe_session_id,omitempty"` - CheckoutURL string `json:"checkout_url,omitempty"` - CreatedAt time.Time `json:"created_at"` - PaidAt *time.Time `json:"paid_at,omitempty"` + ID string `json:"id"` + TenantID string `json:"tenant_id"` + AmountMinor int64 `json:"amount_minor"` + AmountMicros int64 `json:"amount_micros"` + Currency string `json:"currency"` + Status string `json:"status"` + StripeSessionID string `json:"stripe_session_id,omitempty"` + CheckoutURL string `json:"checkout_url,omitempty"` + CreatedAt time.Time `json:"created_at"` + PaidAt *time.Time `json:"paid_at,omitempty"` + StripeCustomerID string `json:"stripe_customer_id,omitempty"` + StripePaymentIntentID string `json:"stripe_payment_intent_id,omitempty"` + StripeChargeID string `json:"stripe_charge_id,omitempty"` + StripeInvoiceID string `json:"stripe_invoice_id,omitempty"` + InvoiceURL string `json:"invoice_url,omitempty"` + InvoicePDFURL string `json:"invoice_pdf_url,omitempty"` + ReceiptURL string `json:"receipt_url,omitempty"` + RefundedMicros int64 `json:"refunded_micros"` + DisputedMicros int64 `json:"disputed_micros"` + ReconciliationStatus string `json:"reconciliation_status"` + ReconciledAt *time.Time `json:"reconciled_at,omitempty"` + ReconciliationError string `json:"reconciliation_error,omitempty"` +} + +type ResolveMissingTopUpInput struct { + Reason string `json:"reason"` +} + +type ResolutionActor struct { + ID string + Type string +} + +type RefundInput struct { + AmountMinor int64 `json:"amount_minor"` + Reason string `json:"reason"` +} + +type Refund struct { + ID string `json:"id"` + TenantID string `json:"tenant_id"` + TopUpOrderID string `json:"topup_order_id"` + StripeRefundID string `json:"stripe_refund_id,omitempty"` + AmountMinor int64 `json:"amount_minor"` + AmountMicros int64 `json:"amount_micros"` + Currency string `json:"currency"` + Reason string `json:"reason"` + Status string `json:"status"` + LastError string `json:"last_error,omitempty"` + CreatedAt time.Time `json:"created_at"` + CompletedAt *time.Time `json:"completed_at,omitempty"` +} + +type PortalResult struct { + URL string `json:"url"` +} + +type ReconciliationResult struct { + ID string `json:"id"` + Status string `json:"status"` + CheckedOrders int64 `json:"checked_orders"` + MismatchCount int64 `json:"mismatch_count"` + Mismatches []map[string]any `json:"mismatches"` + Repairs []map[string]any `json:"repairs,omitempty"` +} + +type PaymentDispute struct { + ID string `json:"id"` + TenantID string `json:"tenant_id,omitempty"` + TopUpOrderID string `json:"topup_order_id,omitempty"` + AmountMinor int64 `json:"amount_minor"` + AmountMicros int64 `json:"amount_micros"` + Currency string `json:"currency"` + Status string `json:"status"` + Reason string `json:"reason"` + DebitedMicros int64 `json:"debited_micros"` + UncollectedMicros int64 `json:"uncollected_micros"` + DueBy *time.Time `json:"due_by,omitempty"` + UpdatedAt time.Time `json:"updated_at"` +} + +type Invoice struct { + ID string `json:"id"` + TenantID string `json:"tenant_id,omitempty"` + TopUpOrderID string `json:"topup_order_id,omitempty"` + Status string `json:"status"` + Currency string `json:"currency"` + AmountDueMinor int64 `json:"amount_due_minor"` + AmountPaidMinor int64 `json:"amount_paid_minor"` + AttemptCount int `json:"attempt_count"` + NextPaymentAttempt *time.Time `json:"next_payment_attempt,omitempty"` + HostedInvoiceURL string `json:"hosted_invoice_url,omitempty"` + InvoicePDFURL string `json:"invoice_pdf_url,omitempty"` + LastFailure string `json:"last_failure,omitempty"` + UpdatedAt time.Time `json:"updated_at"` } diff --git a/internal/catalog/catalog.go b/internal/catalog/catalog.go index 7966d9a..09a6fba 100644 --- a/internal/catalog/catalog.go +++ b/internal/catalog/catalog.go @@ -15,8 +15,9 @@ type Catalog struct { } type snapshot struct { - models map[string]domain.Model - list []domain.Model + models map[string]domain.Model + aliases map[string]string + list []domain.Model } func New(cfg config.Config) *Catalog { @@ -61,15 +62,19 @@ func NewModels(models []domain.Model) *Catalog { func (c *Catalog) Replace(source []domain.Model) { models := make(map[string]domain.Model, len(source)) + aliases := make(map[string]string) list := make([]domain.Model, 0, len(source)) for _, sourceModel := range source { model := sourceModel model.Routes = append([]domain.Route(nil), sourceModel.Routes...) models[model.ID] = model + for _, alias := range model.Aliases { + aliases[alias] = model.ID + } list = append(list, model) } sort.Slice(list, func(i, j int) bool { return list[i].ID < list[j].ID }) - c.state.Store(&snapshot{models: models, list: list}) + c.state.Store(&snapshot{models: models, aliases: aliases, list: list}) } func (c *Catalog) Model(id string) (domain.Model, error) { @@ -79,11 +84,38 @@ func (c *Catalog) Model(id string) (domain.Model, error) { } model, ok := current.models[id] if !ok { + if canonical, aliasOK := current.aliases[id]; aliasOK { + model, ok = current.models[canonical] + } + } + if !ok { return domain.Model{}, fmt.Errorf("model %q not found", id) } return model, nil } +func (c *Catalog) ModelForPrincipal(id string, principal domain.Principal) (domain.Model, error) { + model, err := c.Model(id) + if err != nil { + return domain.Model{}, err + } + if !model.Allows(principal) { + return domain.Model{}, fmt.Errorf("model %q not allowed", id) + } + return model, nil +} + +func (c *Catalog) ModelsFor(protocol domain.Protocol, principal domain.Principal) []domain.Model { + models := c.Models(protocol) + result := models[:0] + for _, model := range models { + if model.Allows(principal) { + result = append(result, model) + } + } + return result +} + func (c *Catalog) Models(protocol domain.Protocol) []domain.Model { current := c.state.Load() if current == nil { diff --git a/internal/config/config.go b/internal/config/config.go index 376cf77..a67f221 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -5,6 +5,7 @@ import ( "errors" "fmt" "io" + "net" "net/url" "os" "strings" @@ -26,12 +27,27 @@ type Config struct { } type ServerConfig struct { - Address string `json:"-"` - AddressEnv string `json:"address_env"` - MaxBodyBytes int64 `json:"max_body_bytes"` - ReadHeaderTimeoutSecs int `json:"read_header_timeout_seconds"` - IdleTimeoutSecs int `json:"idle_timeout_seconds"` - ShutdownTimeoutSecs int `json:"shutdown_timeout_seconds"` + Address string `json:"-"` + AddressEnv string `json:"address_env"` + SplitListeners bool `json:"split_listeners"` + PublicAddressEnv string `json:"public_address_env"` + AdminAddressEnv string `json:"admin_address_env"` + WebhookAddressEnv string `json:"webhook_address_env"` + OperationsAddressEnv string `json:"operations_address_env"` + TrustedProxyCIDRsEnv string `json:"trusted_proxy_cidrs_env"` + RequireHTTPSEnv string `json:"require_https_env"` + DeploymentRegionEnv string `json:"deployment_region_env"` + PublicAddress string `json:"-"` + AdminAddress string `json:"-"` + WebhookAddress string `json:"-"` + OperationsAddress string `json:"-"` + TrustedProxyCIDRs []string `json:"-"` + RequireHTTPS bool `json:"-"` + DeploymentRegion string `json:"-"` + MaxBodyBytes int64 `json:"max_body_bytes"` + ReadHeaderTimeoutSecs int `json:"read_header_timeout_seconds"` + IdleTimeoutSecs int `json:"idle_timeout_seconds"` + ShutdownTimeoutSecs int `json:"shutdown_timeout_seconds"` } type AuthConfig struct { @@ -40,45 +56,55 @@ type AuthConfig struct { } type ControlPlaneConfig struct { - Enabled bool `json:"enabled"` - DatabaseURLEnv string `json:"database_url_env"` - RedisURLEnv string `json:"redis_url_env"` - CredentialKeyEnv string `json:"credential_key_env"` - RedisChannel string `json:"redis_channel"` - SnapshotCacheKey string `json:"snapshot_cache_key"` - ReloadIntervalSeconds int `json:"reload_interval_seconds"` - AutoMigrate bool `json:"auto_migrate"` - DatabaseURL string `json:"-"` - RedisURL string `json:"-"` - CredentialKey string `json:"-"` + Enabled bool `json:"enabled"` + DatabaseURLEnv string `json:"database_url_env"` + RedisURLEnv string `json:"redis_url_env"` + CredentialKeyEnv string `json:"credential_key_env"` + PreviousCredentialKeysEnv string `json:"previous_credential_keys_env"` + RedisChannel string `json:"redis_channel"` + SnapshotCacheKey string `json:"snapshot_cache_key"` + ReloadIntervalSeconds int `json:"reload_interval_seconds"` + AutoMigrate bool `json:"auto_migrate"` + DatabaseURL string `json:"-"` + RedisURL string `json:"-"` + CredentialKey string `json:"-"` + PreviousCredentialKeys []string `json:"-"` } type AdminConfig struct { - Enabled bool `json:"enabled"` - TokenEnv string `json:"token_env"` - BasePath string `json:"base_path"` - RegistrationEnabled bool `json:"registration_enabled"` - SessionTTLHours int `json:"session_ttl_hours"` - PublicURL string `json:"-"` - PublicURLEnv string `json:"public_url_env"` - Mail MailConfig `json:"mail"` - WebAuthn WebAuthnConfig `json:"webauthn"` - Token string `json:"-"` + Enabled bool `json:"enabled"` + TokenEnv string `json:"token_env"` + BasePath string `json:"base_path"` + RegistrationEnabled bool `json:"registration_enabled"` + SessionTTLHours int `json:"session_ttl_hours"` + AuditRetentionDays int `json:"audit_retention_days"` + SecurityRetentionDays int `json:"security_retention_days"` + PublicURL string `json:"-"` + PublicURLEnv string `json:"public_url_env"` + Mail MailConfig `json:"mail"` + WebAuthn WebAuthnConfig `json:"webauthn"` + Token string `json:"-"` } type MailConfig struct { - Enabled bool `json:"enabled"` - FromName string `json:"from_name"` - TLSMode string `json:"tls_mode"` - FromAddressEnv string `json:"from_address_env"` - SMTPAddressEnv string `json:"smtp_address_env"` - SMTPUsernameEnv string `json:"smtp_username_env"` - SMTPPasswordEnv string `json:"smtp_password_env"` - SMTPImplicitTLS bool `json:"smtp_implicit_tls"` - FromAddress string `json:"-"` - SMTPAddress string `json:"-"` - SMTPUsername string `json:"-"` - SMTPPassword string `json:"-"` + Enabled bool `json:"enabled"` + FromName string `json:"from_name"` + TLSMode string `json:"tls_mode"` + FromAddressEnv string `json:"from_address_env"` + SMTPAddressEnv string `json:"smtp_address_env"` + SMTPUsernameEnv string `json:"smtp_username_env"` + SMTPPasswordEnv string `json:"smtp_password_env"` + FeedbackSecretEnv string `json:"feedback_secret_env"` + LowBalanceMicros int64 `json:"low_balance_micros"` + SpendAnomalyMultiplier int64 `json:"spend_anomaly_multiplier"` + SpendAnomalyMinMicros int64 `json:"spend_anomaly_min_micros"` + NotificationIntervalSeconds int `json:"notification_interval_seconds"` + SMTPImplicitTLS bool `json:"smtp_implicit_tls"` + FromAddress string `json:"-"` + SMTPAddress string `json:"-"` + SMTPUsername string `json:"-"` + SMTPPassword string `json:"-"` + FeedbackSecret string `json:"-"` } type WebAuthnConfig struct { @@ -134,19 +160,29 @@ type BillingConfig struct { DefaultMaxOutputTokens int64 `json:"default_max_output_tokens"` MinTopUpMinor int64 `json:"min_top_up_minor"` MaxTopUpMinor int64 `json:"max_top_up_minor"` + SettlementSpoolPathEnv string `json:"settlement_spool_path_env"` + SettlementSpoolPath string `json:"-"` Stripe StripeConfig `json:"stripe"` } type StripeConfig struct { - Enabled bool `json:"enabled"` - APIKeyEnv string `json:"api_key_env"` - WebhookSecretEnv string `json:"webhook_secret_env"` - SuccessURL string `json:"-"` - CancelURL string `json:"-"` - SuccessURLEnv string `json:"success_url_env"` - CancelURLEnv string `json:"cancel_url_env"` - APIKey string `json:"-"` - WebhookSecret string `json:"-"` + Enabled bool `json:"enabled"` + APIKeyEnv string `json:"api_key_env"` + WebhookSecretEnv string `json:"webhook_secret_env"` + SuccessURLEnv string `json:"success_url_env"` + CancelURLEnv string `json:"cancel_url_env"` + PortalReturnURLEnv string `json:"portal_return_url_env"` + AutomaticTaxEnabledEnv string `json:"automatic_tax_enabled_env"` + TaxRegistrationConfirmedEnv string `json:"tax_registration_confirmed_env"` + ProductTaxCodeEnv string `json:"product_tax_code_env"` + SuccessURL string `json:"-"` + CancelURL string `json:"-"` + PortalReturnURL string `json:"-"` + APIKey string `json:"-"` + WebhookSecret string `json:"-"` + AutomaticTaxEnabled bool `json:"-"` + TaxRegistrationConfirmed bool `json:"-"` + ProductTaxCode string `json:"-"` } func Load(path string) (Config, error) { @@ -186,6 +222,27 @@ func applyDefaults(cfg *Config) { if cfg.Server.Address == "" { cfg.Server.Address = ":8080" } + if cfg.Server.PublicAddressEnv == "" { + cfg.Server.PublicAddressEnv = "AIGW_PUBLIC_ADDRESS" + } + if cfg.Server.AdminAddressEnv == "" { + cfg.Server.AdminAddressEnv = "AIGW_ADMIN_ADDRESS" + } + if cfg.Server.WebhookAddressEnv == "" { + cfg.Server.WebhookAddressEnv = "AIGW_WEBHOOK_ADDRESS" + } + if cfg.Server.OperationsAddressEnv == "" { + cfg.Server.OperationsAddressEnv = "AIGW_OPERATIONS_ADDRESS" + } + if cfg.Server.TrustedProxyCIDRsEnv == "" { + cfg.Server.TrustedProxyCIDRsEnv = "AIGW_TRUSTED_PROXY_CIDRS" + } + if cfg.Server.RequireHTTPSEnv == "" { + cfg.Server.RequireHTTPSEnv = "AIGW_REQUIRE_HTTPS" + } + if cfg.Server.DeploymentRegionEnv == "" { + cfg.Server.DeploymentRegionEnv = "AIGW_DEPLOYMENT_REGION" + } if cfg.Server.MaxBodyBytes == 0 { cfg.Server.MaxBodyBytes = 16 << 20 } @@ -210,6 +267,9 @@ func applyDefaults(cfg *Config) { if cfg.ControlPlane.CredentialKeyEnv == "" { cfg.ControlPlane.CredentialKeyEnv = "AIGW_CREDENTIAL_KEY" } + if cfg.ControlPlane.PreviousCredentialKeysEnv == "" { + cfg.ControlPlane.PreviousCredentialKeysEnv = "AIGW_CREDENTIAL_PREVIOUS_KEYS" + } if cfg.ControlPlane.RedisChannel == "" { cfg.ControlPlane.RedisChannel = "aigw:control:changed" } @@ -228,6 +288,12 @@ func applyDefaults(cfg *Config) { if cfg.Admin.SessionTTLHours == 0 { cfg.Admin.SessionTTLHours = 12 } + if cfg.Admin.AuditRetentionDays == 0 { + cfg.Admin.AuditRetentionDays = 2555 + } + if cfg.Admin.SecurityRetentionDays == 0 { + cfg.Admin.SecurityRetentionDays = 30 + } if cfg.Admin.PublicURLEnv == "" { cfg.Admin.PublicURLEnv = "AIGW_PUBLIC_URL" } @@ -249,6 +315,21 @@ func applyDefaults(cfg *Config) { if cfg.Admin.Mail.SMTPPasswordEnv == "" { cfg.Admin.Mail.SMTPPasswordEnv = "AIGW_SMTP_PASSWORD" } + if cfg.Admin.Mail.FeedbackSecretEnv == "" { + cfg.Admin.Mail.FeedbackSecretEnv = "AIGW_MAIL_FEEDBACK_SECRET" + } + if cfg.Admin.Mail.LowBalanceMicros == 0 { + cfg.Admin.Mail.LowBalanceMicros = 5_000_000 + } + if cfg.Admin.Mail.SpendAnomalyMultiplier == 0 { + cfg.Admin.Mail.SpendAnomalyMultiplier = 3 + } + if cfg.Admin.Mail.SpendAnomalyMinMicros == 0 { + cfg.Admin.Mail.SpendAnomalyMinMicros = 10_000_000 + } + if cfg.Admin.Mail.NotificationIntervalSeconds == 0 { + cfg.Admin.Mail.NotificationIntervalSeconds = 300 + } if cfg.Admin.WebAuthn.RPDisplayName == "" { cfg.Admin.WebAuthn.RPDisplayName = "AIGW Console" } @@ -285,6 +366,9 @@ func applyDefaults(cfg *Config) { if cfg.Billing.MinTopUpMinor == 0 { cfg.Billing.MinTopUpMinor = 500 } + if cfg.Billing.SettlementSpoolPathEnv == "" { + cfg.Billing.SettlementSpoolPathEnv = "AIGW_SETTLEMENT_SPOOL_PATH" + } if cfg.Billing.Stripe.APIKeyEnv == "" { cfg.Billing.Stripe.APIKeyEnv = "AIGW_STRIPE_API_KEY" } @@ -297,6 +381,18 @@ func applyDefaults(cfg *Config) { if cfg.Billing.Stripe.CancelURLEnv == "" { cfg.Billing.Stripe.CancelURLEnv = "AIGW_STRIPE_CANCEL_URL" } + if cfg.Billing.Stripe.PortalReturnURLEnv == "" { + cfg.Billing.Stripe.PortalReturnURLEnv = "AIGW_STRIPE_PORTAL_RETURN_URL" + } + if cfg.Billing.Stripe.AutomaticTaxEnabledEnv == "" { + cfg.Billing.Stripe.AutomaticTaxEnabledEnv = "AIGW_STRIPE_AUTOMATIC_TAX_ENABLED" + } + if cfg.Billing.Stripe.TaxRegistrationConfirmedEnv == "" { + cfg.Billing.Stripe.TaxRegistrationConfirmedEnv = "AIGW_STRIPE_TAX_REGISTRATION_CONFIRMED" + } + if cfg.Billing.Stripe.ProductTaxCodeEnv == "" { + cfg.Billing.Stripe.ProductTaxCodeEnv = "AIGW_STRIPE_PRODUCT_TAX_CODE" + } for i := range cfg.Models { for j := range cfg.Models[i].Routes { if cfg.Models[i].Routes[j].Weight == 0 { @@ -310,10 +406,40 @@ func resolveSecrets(cfg *Config) error { if value := strings.TrimSpace(os.Getenv(cfg.Server.AddressEnv)); value != "" { cfg.Server.Address = value } + if cfg.Server.SplitListeners { + if err := resolveRequiredEnv(&cfg.Server.PublicAddress, cfg.Server.PublicAddressEnv, "server.public_address"); err != nil { + return err + } + if err := resolveRequiredEnv(&cfg.Server.AdminAddress, cfg.Server.AdminAddressEnv, "server.admin_address"); err != nil { + return err + } + if err := resolveRequiredEnv(&cfg.Server.WebhookAddress, cfg.Server.WebhookAddressEnv, "server.webhook_address"); err != nil { + return err + } + if err := resolveRequiredEnv(&cfg.Server.OperationsAddress, cfg.Server.OperationsAddressEnv, "server.operations_address"); err != nil { + return err + } + } + for _, cidr := range strings.Split(os.Getenv(cfg.Server.TrustedProxyCIDRsEnv), ",") { + if cidr = strings.TrimSpace(cidr); cidr != "" { + cfg.Server.TrustedProxyCIDRs = append(cfg.Server.TrustedProxyCIDRs, cidr) + } + } + var err error + cfg.Server.RequireHTTPS, err = envBool(cfg.Server.RequireHTTPSEnv) + if err != nil { + return err + } + cfg.Server.DeploymentRegion = strings.ToLower(strings.TrimSpace(os.Getenv(cfg.Server.DeploymentRegionEnv))) if cfg.ControlPlane.Enabled { cfg.ControlPlane.DatabaseURL = os.Getenv(cfg.ControlPlane.DatabaseURLEnv) cfg.ControlPlane.RedisURL = os.Getenv(cfg.ControlPlane.RedisURLEnv) cfg.ControlPlane.CredentialKey = os.Getenv(cfg.ControlPlane.CredentialKeyEnv) + for _, value := range strings.Split(os.Getenv(cfg.ControlPlane.PreviousCredentialKeysEnv), ",") { + if value = strings.TrimSpace(value); value != "" { + cfg.ControlPlane.PreviousCredentialKeys = append(cfg.ControlPlane.PreviousCredentialKeys, value) + } + } } if cfg.Admin.Enabled { cfg.Admin.Token = os.Getenv(cfg.Admin.TokenEnv) @@ -325,6 +451,7 @@ func resolveSecrets(cfg *Config) error { cfg.Admin.Mail.SMTPAddress = strings.TrimSpace(os.Getenv(cfg.Admin.Mail.SMTPAddressEnv)) cfg.Admin.Mail.SMTPUsername = os.Getenv(cfg.Admin.Mail.SMTPUsernameEnv) cfg.Admin.Mail.SMTPPassword = os.Getenv(cfg.Admin.Mail.SMTPPasswordEnv) + cfg.Admin.Mail.FeedbackSecret = os.Getenv(cfg.Admin.Mail.FeedbackSecretEnv) } if cfg.Admin.WebAuthn.Enabled { cfg.Admin.WebAuthn.RPID = strings.TrimSpace(os.Getenv(cfg.Admin.WebAuthn.RPIDEnv)) @@ -344,6 +471,22 @@ func resolveSecrets(cfg *Config) error { if err := resolveRequiredEnv(&cfg.Billing.Stripe.CancelURL, cfg.Billing.Stripe.CancelURLEnv, "billing.stripe.cancel_url"); err != nil { return err } + if err := resolveRequiredEnv(&cfg.Billing.Stripe.PortalReturnURL, cfg.Billing.Stripe.PortalReturnURLEnv, "billing.stripe.portal_return_url"); err != nil { + return err + } + var err error + cfg.Billing.Stripe.AutomaticTaxEnabled, err = envBool(cfg.Billing.Stripe.AutomaticTaxEnabledEnv) + if err != nil { + return err + } + cfg.Billing.Stripe.TaxRegistrationConfirmed, err = envBool(cfg.Billing.Stripe.TaxRegistrationConfirmedEnv) + if err != nil { + return err + } + cfg.Billing.Stripe.ProductTaxCode = strings.TrimSpace(os.Getenv(cfg.Billing.Stripe.ProductTaxCodeEnv)) + } + if cfg.Billing.Enabled { + cfg.Billing.SettlementSpoolPath = strings.TrimSpace(os.Getenv(cfg.Billing.SettlementSpoolPathEnv)) } for i := range cfg.Providers { provider := &cfg.Providers[i] @@ -364,6 +507,17 @@ func resolveSecrets(cfg *Config) error { return nil } +func envBool(name string) (bool, error) { + value := strings.TrimSpace(os.Getenv(name)) + if value == "" || strings.EqualFold(value, "false") || value == "0" { + return false, nil + } + if strings.EqualFold(value, "true") || value == "1" { + return true, nil + } + return false, fmt.Errorf("environment variable %s must be true/false or 1/0", name) +} + func resolveRequiredEnv(target *string, environment, field string) error { if environment == "" { return nil @@ -383,6 +537,23 @@ func Validate(cfg Config) error { if cfg.Observability.UsageBuffer < 1 { return errors.New("observability.usage_buffer must be positive") } + if cfg.Server.SplitListeners { + seen := map[string]string{} + for name, address := range map[string]string{"public": cfg.Server.PublicAddress, "admin": cfg.Server.AdminAddress, "webhook": cfg.Server.WebhookAddress, "operations": cfg.Server.OperationsAddress} { + if strings.TrimSpace(address) == "" { + return fmt.Errorf("server %s listener address is empty", name) + } + if previous, ok := seen[address]; ok { + return fmt.Errorf("server %s and %s listeners must use different addresses", previous, name) + } + seen[address] = name + } + } + for _, cidr := range cfg.Server.TrustedProxyCIDRs { + if _, _, err := net.ParseCIDR(cidr); err != nil { + return fmt.Errorf("invalid trusted proxy CIDR %q", cidr) + } + } if cfg.ControlPlane.Enabled { if cfg.ControlPlane.DatabaseURL == "" { @@ -408,6 +579,9 @@ func Validate(cfg Config) error { if cfg.Admin.SessionTTLHours < 1 || cfg.Admin.SessionTTLHours > 720 { return errors.New("admin.session_ttl_hours must be between 1 and 720") } + if cfg.Admin.AuditRetentionDays < 30 || cfg.Admin.SecurityRetentionDays < 1 { + return errors.New("admin audit retention must be at least 30 days and security retention at least 1 day") + } publicURL, err := url.Parse(cfg.Admin.PublicURL) if err != nil || publicURL.Host == "" || (publicURL.Scheme != "http" && publicURL.Scheme != "https") { return errors.New("admin.public_url must resolve from an environment variable to an absolute http(s) URL") @@ -419,6 +593,13 @@ func Validate(cfg Config) error { if (cfg.Admin.Mail.SMTPUsername == "") != (cfg.Admin.Mail.SMTPPassword == "") { return errors.New("admin.mail SMTP username and password must both be set or both be empty") } + if cfg.Admin.Mail.FeedbackSecret != "" && len(cfg.Admin.Mail.FeedbackSecret) < 32 { + return errors.New("admin.mail feedback secret must contain at least 32 characters") + } + if cfg.Admin.Mail.LowBalanceMicros < 0 || cfg.Admin.Mail.SpendAnomalyMultiplier < 2 || + cfg.Admin.Mail.SpendAnomalyMinMicros < 0 || cfg.Admin.Mail.NotificationIntervalSeconds < 30 { + return errors.New("admin.mail notification thresholds are invalid") + } switch cfg.Admin.Mail.TLSMode { case "starttls", "tls", "none": default: @@ -453,6 +634,9 @@ func Validate(cfg Config) error { if cfg.Billing.MinTopUpMinor < 1 || cfg.Billing.MaxTopUpMinor < cfg.Billing.MinTopUpMinor { return errors.New("billing top-up bounds are invalid") } + if strings.TrimSpace(cfg.Billing.SettlementSpoolPath) == "" { + return fmt.Errorf("billing: environment variable %s is empty; durable settlement fallback is required", cfg.Billing.SettlementSpoolPathEnv) + } if cfg.Billing.Stripe.Enabled { if cfg.Billing.Stripe.APIKey == "" { return fmt.Errorf("billing.stripe: environment variable %s is empty", cfg.Billing.Stripe.APIKeyEnv) @@ -460,14 +644,17 @@ func Validate(cfg Config) error { if cfg.Billing.Stripe.WebhookSecret == "" { return fmt.Errorf("billing.stripe: environment variable %s is empty", cfg.Billing.Stripe.WebhookSecretEnv) } - if strings.TrimSpace(cfg.Billing.Stripe.SuccessURL) == "" || strings.TrimSpace(cfg.Billing.Stripe.CancelURL) == "" { - return errors.New("billing.stripe.success_url and cancel_url are required") + if strings.TrimSpace(cfg.Billing.Stripe.SuccessURL) == "" || strings.TrimSpace(cfg.Billing.Stripe.CancelURL) == "" || strings.TrimSpace(cfg.Billing.Stripe.PortalReturnURL) == "" { + return errors.New("billing.stripe success, cancel, and portal return URLs are required") } - for name, value := range map[string]string{"success_url": cfg.Billing.Stripe.SuccessURL, "cancel_url": cfg.Billing.Stripe.CancelURL} { + for name, value := range map[string]string{"success_url": cfg.Billing.Stripe.SuccessURL, "cancel_url": cfg.Billing.Stripe.CancelURL, "portal_return_url": cfg.Billing.Stripe.PortalReturnURL} { parsed, err := url.Parse(value) if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") { return fmt.Errorf("billing.stripe.%s must be an absolute http(s) URL", name) } + if cfg.Billing.Stripe.AutomaticTaxEnabled && (!cfg.Billing.Stripe.TaxRegistrationConfirmed || cfg.Billing.Stripe.ProductTaxCode == "") { + return errors.New("Stripe automatic tax requires an explicit confirmed registration and product tax code") + } } } } diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 1ecf70f..8e2e9bb 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -110,6 +110,8 @@ func TestLoadResolvesStripeSecrets(t *testing.T) { t.Setenv("TEST_STRIPE_WEBHOOK", "whsec_example") t.Setenv("TEST_STRIPE_SUCCESS", "https://console.example.test/admin/?topup=success") t.Setenv("TEST_STRIPE_CANCEL", "https://console.example.test/admin/?topup=cancel") + t.Setenv("AIGW_STRIPE_PORTAL_RETURN_URL", "https://console.example.test/admin/?billing=portal") + t.Setenv("AIGW_SETTLEMENT_SPOOL_PATH", filepath.Join(t.TempDir(), "settlements.jsonl")) path := writeConfig(t, `{ "control_plane":{"enabled":true}, "billing":{"enabled":true,"stripe":{"enabled":true,"api_key_env":"TEST_STRIPE_KEY","webhook_secret_env":"TEST_STRIPE_WEBHOOK","success_url_env":"TEST_STRIPE_SUCCESS","cancel_url_env":"TEST_STRIPE_CANCEL"}} @@ -147,6 +149,10 @@ func TestLoadRejectsBillingWithoutControlPlane(t *testing.T) { func TestVersionedExamplesResolveExternalServicesFromEnvironment(t *testing.T) { t.Setenv("AIGW_SERVER_ADDRESS", "127.0.0.1:18081") + t.Setenv("AIGW_PUBLIC_ADDRESS", "127.0.0.1:18081") + t.Setenv("AIGW_ADMIN_ADDRESS", "127.0.0.1:18082") + t.Setenv("AIGW_WEBHOOK_ADDRESS", "127.0.0.1:18083") + t.Setenv("AIGW_OPERATIONS_ADDRESS", "127.0.0.1:19090") t.Setenv("AIGW_API_KEYS", `[{"key":"test","key_id":"key","tenant_id":"tenant","project_id":"project","scopes":["inference"]}]`) t.Setenv("OPENAI_BASE_URL", "https://openai.example.test/v1") t.Setenv("OPENAI_API_KEY", "openai-secret") @@ -175,6 +181,8 @@ func TestVersionedExamplesResolveExternalServicesFromEnvironment(t *testing.T) { t.Setenv("AIGW_STRIPE_WEBHOOK_SECRET", "whsec_example") t.Setenv("AIGW_STRIPE_SUCCESS_URL", "https://console.example.test/admin/?topup=success") t.Setenv("AIGW_STRIPE_CANCEL_URL", "https://console.example.test/admin/?topup=cancel") + t.Setenv("AIGW_STRIPE_PORTAL_RETURN_URL", "https://console.example.test/admin/?billing=portal") + t.Setenv("AIGW_SETTLEMENT_SPOOL_PATH", filepath.Join(t.TempDir(), "settlements.jsonl")) controlConfig, err := Load(filepath.Join("..", "..", "config.control.example.json")) if err != nil { t.Fatalf("load control-plane example: %v", err) diff --git a/internal/controlplane/mail_operations.go b/internal/controlplane/mail_operations.go new file mode 100644 index 0000000..aaaf2c5 --- /dev/null +++ b/internal/controlplane/mail_operations.go @@ -0,0 +1,213 @@ +package controlplane + +import ( + "context" + "crypto/hmac" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "log/slog" + "net/http" + "strconv" + "strings" + "time" + + "github.com/jackc/pgx/v5" +) + +const maxMailFeedbackBytes = 64 << 10 + +type MailNotificationConfig struct { + LowBalanceMicros int64 + SpendAnomalyMultiplier int64 + SpendAnomalyMinMicros int64 + Interval time.Duration +} + +type mailFeedback struct { + EventID string `json:"event_id"` + EventType string `json:"event_type"` + Recipient string `json:"recipient"` + Provider string `json:"provider"` + Detail string `json:"detail"` +} + +// MailFeedbackHandler accepts a provider-neutral normalized callback. An edge +// adapter maps the provider's native event into this payload and signs +// "<unix timestamp>.<raw body>" with HMAC-SHA256. +func (s *Store) MailFeedbackHandler(secret string) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + w.Header().Set("Allow", http.MethodPost) + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + return + } + body, err := io.ReadAll(io.LimitReader(r.Body, maxMailFeedbackBytes+1)) + if err != nil || len(body) > maxMailFeedbackBytes || !validMailFeedbackSignature(secret, r.Header.Get("X-AIGW-Mail-Timestamp"), r.Header.Get("X-AIGW-Mail-Signature"), body, time.Now().UTC()) { + http.Error(w, "invalid feedback signature", http.StatusBadRequest) + return + } + var event mailFeedback + decoder := json.NewDecoder(strings.NewReader(string(body))) + decoder.DisallowUnknownFields() + if decoder.Decode(&event) != nil || strings.TrimSpace(event.EventID) == "" || + !strings.Contains(event.Recipient, "@") || (event.EventType != "delivered" && event.EventType != "bounce" && event.EventType != "complaint") { + http.Error(w, "invalid feedback event", http.StatusBadRequest) + return + } + if err := s.applyMailFeedback(r.Context(), event); err != nil { + http.Error(w, "feedback persistence failed", http.StatusInternalServerError) + return + } + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, "{\"received\":true}\n") + }) +} + +func validMailFeedbackSignature(secret, timestamp, signature string, body []byte, now time.Time) bool { + if secret == "" || timestamp == "" || !strings.HasPrefix(signature, "sha256=") { + return false + } + seconds, err := strconv.ParseInt(timestamp, 10, 64) + if err != nil || now.Sub(time.Unix(seconds, 0)).Abs() > 5*time.Minute { + return false + } + provided, err := hex.DecodeString(strings.TrimPrefix(signature, "sha256=")) + if err != nil { + return false + } + mac := hmac.New(sha256.New, []byte(secret)) + _, _ = mac.Write([]byte(timestamp)) + _, _ = mac.Write([]byte(".")) + _, _ = mac.Write(body) + return hmac.Equal(provided, mac.Sum(nil)) +} + +func (s *Store) applyMailFeedback(ctx context.Context, event mailFeedback) error { + tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}) + if err != nil { + return err + } + defer tx.Rollback(ctx) + tag, err := tx.Exec(ctx, `INSERT INTO mail_feedback_events (event_id,event_type,recipient,provider) + VALUES ($1,$2,lower($3),$4) ON CONFLICT DO NOTHING`, event.EventID, event.EventType, event.Recipient, event.Provider) + if err != nil || tag.RowsAffected() == 0 { + if err != nil { + return err + } + return tx.Commit(ctx) + } + if event.EventType == "bounce" || event.EventType == "complaint" { + detail := event.Detail + if len(detail) > 1000 { + detail = detail[:1000] + } + if _, err := tx.Exec(ctx, `INSERT INTO mail_suppressions + (recipient,reason,provider,provider_event_id,detail) VALUES (lower($1),$2,$3,$4,$5) + ON CONFLICT (recipient) DO UPDATE SET reason=EXCLUDED.reason,provider=EXCLUDED.provider, + provider_event_id=EXCLUDED.provider_event_id,detail=EXCLUDED.detail,updated_at=now()`, + event.Recipient, event.EventType, event.Provider, event.EventID, detail); err != nil { + return err + } + if _, err := tx.Exec(ctx, `UPDATE console_mail_outbox SET status='suppressed',last_error=$2,claimed_at=NULL + WHERE lower(recipient)=lower($1) AND status IN ('pending','retry','sending')`, event.Recipient, "recipient suppressed after "+event.EventType); err != nil { + return err + } + } + return tx.Commit(ctx) +} + +func (s *Store) RunMailNotificationWorker(ctx context.Context, config MailNotificationConfig, logger *slog.Logger) { + if config.Interval <= 0 { + config.Interval = 5 * time.Minute + } + if logger == nil { + logger = slog.Default() + } + ticker := time.NewTicker(config.Interval) + defer ticker.Stop() + for { + if err := s.queueBillingNotifications(ctx, config); err != nil && !errors.Is(err, context.Canceled) { + logger.Warn("billing_notification_scan_failed", "error", err) + } + select { + case <-ctx.Done(): + return + case <-ticker.C: + } + } +} + +func (s *Store) queueBillingNotifications(ctx context.Context, config MailNotificationConfig) error { + rows, err := s.db.Query(ctx, `WITH spend AS ( + SELECT tenant_id, + COALESCE(sum(-amount_micros) FILTER (WHERE kind='usage' AND created_at>=date_trunc('day',now())),0)::bigint today, + (COALESCE(sum(-amount_micros) FILTER (WHERE kind='usage' AND created_at>=date_trunc('day',now())-interval '7 days' AND created_at<date_trunc('day',now())),0)/7)::bigint baseline + FROM billing_ledger GROUP BY tenant_id) + SELECT w.tenant_id::text,w.currency,w.balance_micros-w.reserved_micros,u.email,u.display_name, + COALESCE(spend.today,0),COALESCE(spend.baseline,0) + FROM tenant_wallets w JOIN console_users u ON u.tenant_id=w.tenant_id + LEFT JOIN spend ON spend.tenant_id=w.tenant_id + WHERE u.status='active' AND u.email_verified_at IS NOT NULL AND u.role IN ('tenant_admin','tenant_billing') + AND NOT EXISTS (SELECT 1 FROM mail_suppressions s WHERE lower(s.recipient)=lower(u.email))`) + if err != nil { + return err + } + defer rows.Close() + for rows.Next() { + var tenantID, currency, email, name string + var available, today, baseline int64 + if err := rows.Scan(&tenantID, ¤cy, &available, &email, &name, &today, &baseline); err != nil { + return err + } + day := time.Now().UTC().Format("2006-01-02") + if available <= config.LowBalanceMicros { + body := fmt.Sprintf("Hi %s,\n\nYour AIGW prepaid balance is low: %.6f %s remains available. Add funds to avoid interrupted API access.\n", displayName(name), float64(available)/1_000_000, strings.ToUpper(currency)) + if err := s.queueNotification(ctx, tenantID, email, "low_balance", day, "AIGW balance is low", body); err != nil { + return err + } + } + if baseline > 0 && today >= config.SpendAnomalyMinMicros && today >= baseline*config.SpendAnomalyMultiplier { + body := fmt.Sprintf("Hi %s,\n\nAIGW detected unusual API spend today: %.6f %s versus a seven-day daily baseline of %.6f %s. Review API keys and usage in the console.\n", displayName(name), float64(today)/1_000_000, strings.ToUpper(currency), float64(baseline)/1_000_000, strings.ToUpper(currency)) + if err := s.queueNotification(ctx, tenantID, email, "spend_anomaly", day, "Unusual AIGW API spend detected", body); err != nil { + return err + } + } + } + return rows.Err() +} + +func (s *Store) queueNotification(ctx context.Context, tenantID, recipient, kind, dedupe, subject, body string) error { + tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}) + if err != nil { + return err + } + defer tx.Rollback(ctx) + tag, err := tx.Exec(ctx, `INSERT INTO mail_notification_events (tenant_id,recipient,notification_type,dedupe_key) + VALUES ($1,lower($2),$3,$4) ON CONFLICT DO NOTHING`, tenantID, recipient, kind, dedupe) + if err != nil || tag.RowsAffected() == 0 { + if err != nil { + return err + } + return tx.Commit(ctx) + } + ciphertext, err := s.cipher.Encrypt(body) + if err != nil { + return err + } + if _, err := tx.Exec(ctx, `INSERT INTO console_mail_outbox (recipient,template,subject,body_ciphertext) + VALUES (lower($1),$2,$3,$4)`, recipient, kind, subject, ciphertext); err != nil { + return err + } + return tx.Commit(ctx) +} + +func displayName(value string) string { + if value = strings.TrimSpace(value); value != "" { + return value + } + return "there" +} diff --git a/internal/controlplane/mail_operations_test.go b/internal/controlplane/mail_operations_test.go new file mode 100644 index 0000000..0b47cdb --- /dev/null +++ b/internal/controlplane/mail_operations_test.go @@ -0,0 +1,29 @@ +package controlplane + +import ( + "crypto/hmac" + "crypto/sha256" + "encoding/hex" + "strconv" + "testing" + "time" +) + +func TestMailFeedbackSignature(t *testing.T) { + now := time.Unix(1_800_000_000, 0).UTC() + timestamp := strconv.FormatInt(now.Unix(), 10) + body := []byte(`{"event_id":"evt_1","event_type":"bounce","recipient":"test@example.com","provider":"test"}`) + mac := hmac.New(sha256.New, []byte("a-production-length-feedback-secret")) + _, _ = mac.Write([]byte(timestamp + ".")) + _, _ = mac.Write(body) + signature := "sha256=" + hex.EncodeToString(mac.Sum(nil)) + if !validMailFeedbackSignature("a-production-length-feedback-secret", timestamp, signature, body, now) { + t.Fatal("valid signature was rejected") + } + if validMailFeedbackSignature("a-production-length-feedback-secret", timestamp, signature, []byte(`{}`), now) { + t.Fatal("signature must bind the raw body") + } + if validMailFeedbackSignature("a-production-length-feedback-secret", timestamp, signature, body, now.Add(6*time.Minute)) { + t.Fatal("stale signature was accepted") + } +} diff --git a/internal/controlplane/manager.go b/internal/controlplane/manager.go index b6be748..212963b 100644 --- a/internal/controlplane/manager.go +++ b/internal/controlplane/manager.go @@ -28,16 +28,19 @@ type policyReplacer interface { } type Manager struct { - store managerStore - catalog *catalog.Catalog - authenticator *auth.StaticAuthenticator - logger *slog.Logger - pollInterval time.Duration - generation atomic.Int64 - redisConnected atomic.Bool - reloadMu sync.Mutex - broadcasts chan ChangeEvent - policyTarget policyReplacer + store managerStore + catalog *catalog.Catalog + authenticator *auth.StaticAuthenticator + logger *slog.Logger + pollInterval time.Duration + generation atomic.Int64 + redisConnected atomic.Bool + loaded atomic.Bool + lastReloadUnix atomic.Int64 + lastHealthyUnix atomic.Int64 + reloadMu sync.Mutex + broadcasts chan ChangeEvent + policyTarget policyReplacer } func NewManager(store managerStore, modelCatalog *catalog.Catalog, authenticator *auth.StaticAuthenticator, logger *slog.Logger, pollInterval time.Duration, policyTargets ...policyReplacer) *Manager { @@ -67,10 +70,30 @@ func (m *Manager) Reload(ctx context.Context) (int64, error) { m.policyTarget.ReplacePolicies(snapshot.Limits) } m.generation.Store(snapshot.Generation) + m.loaded.Store(true) + m.lastReloadUnix.Store(time.Now().UTC().Unix()) + m.lastHealthyUnix.Store(time.Now().UTC().Unix()) m.logger.Info("control_plane_reloaded", "generation", snapshot.Generation, "models", len(snapshot.Models), "api_keys", len(snapshot.APIKeys)) return snapshot.Generation, nil } +func (m *Manager) Loaded() bool { return m.loaded.Load() } +func (m *Manager) LastReloadAt() time.Time { + value := m.lastReloadUnix.Load() + if value == 0 { + return time.Time{} + } + return time.Unix(value, 0).UTC() +} + +func (m *Manager) LastHealthyAt() time.Time { + value := m.lastHealthyUnix.Load() + if value == 0 { + return time.Time{} + } + return time.Unix(value, 0).UTC() +} + func (m *Manager) AfterMutation(ctx context.Context, generation int64, resource, id string) error { loadedGeneration, err := m.Reload(ctx) if err != nil { @@ -135,6 +158,7 @@ func (m *Manager) runPolling(ctx context.Context) { m.logger.Warn("control_plane_generation_check_failed", "error", err) continue } + m.lastHealthyUnix.Store(time.Now().UTC().Unix()) if generation > m.generation.Load() { if _, err := m.Reload(ctx); err != nil { m.logger.Error("control_plane_reload_failed", "source", "postgres", "error", err) diff --git a/internal/controlplane/mutations.go b/internal/controlplane/mutations.go index 3cedf70..c2cb7d8 100644 --- a/internal/controlplane/mutations.go +++ b/internal/controlplane/mutations.go @@ -11,6 +11,7 @@ import ( "net/url" "regexp" "strings" + "time" "github.com/jackc/pgx/v5" ) @@ -172,10 +173,45 @@ func (s *Store) SetProviderEnabled(ctx context.Context, id string, enabled bool) func (s *Store) CreateModel(ctx context.Context, input CreateModelInput) (Model, int64, error) { input.PublicID = strings.TrimSpace(input.PublicID) + input.DisplayName = strings.TrimSpace(input.DisplayName) + input.Description = strings.TrimSpace(input.Description) input.OwnedBy = strings.TrimSpace(input.OwnedBy) + input.Lifecycle = strings.ToLower(strings.TrimSpace(input.Lifecycle)) + input.PriceCurrency = strings.ToLower(strings.TrimSpace(input.PriceCurrency)) + if input.DisplayName == "" { + input.DisplayName = input.PublicID + } + if input.Lifecycle == "" { + input.Lifecycle = "active" + } + if input.PriceCurrency == "" { + input.PriceCurrency = "usd" + } + if len(input.InputModalities) == 0 { + input.InputModalities = []string{"text"} + } + if len(input.OutputModalities) == 0 { + input.OutputModalities = []string{"text"} + } + if len(input.Capabilities) == 0 { + input.Capabilities = []string{"chat", "streaming"} + } + input.InputModalities = uniqueStrings(input.InputModalities) + input.OutputModalities = uniqueStrings(input.OutputModalities) + input.Capabilities = uniqueStrings(input.Capabilities) + input.Regions = uniqueStrings(input.Regions) + input.Aliases = uniqueStrings(input.Aliases) + input.AllowedTenantIDs = uniqueStrings(input.AllowedTenantIDs) + input.AllowedKeyIDs = uniqueStrings(input.AllowedKeyIDs) if input.PublicID == "" || len(input.Routes) == 0 { return Model{}, 0, errors.New("model requires public_id and at least one route") } + if input.ContextWindow < 0 || input.MaxOutputTokens < 0 || len(input.PriceCurrency) != 3 { + return Model{}, 0, errors.New("model context, output limit, or price currency is invalid") + } + if input.Lifecycle != "preview" && input.Lifecycle != "active" && input.Lifecycle != "deprecated" && input.Lifecycle != "retired" { + return Model{}, 0, errors.New("model lifecycle must be preview, active, deprecated, or retired") + } if input.InputPriceMicrosPerMillion < 0 || input.OutputPriceMicrosPerMillion < 0 || input.CacheReadPriceMicrosPerMillion < 0 || input.CacheWritePriceMicrosPerMillion < 0 { return Model{}, 0, errors.New("model prices cannot be negative") } @@ -194,21 +230,63 @@ func (s *Store) CreateModel(ctx context.Context, input CreateModelInput) (Model, return Model{}, 0, err } defer tx.Rollback(ctx) + inputModalitiesJSON, _ := json.Marshal(input.InputModalities) + outputModalitiesJSON, _ := json.Marshal(input.OutputModalities) + capabilitiesJSON, _ := json.Marshal(input.Capabilities) + regionsJSON, _ := json.Marshal(input.Regions) var result Model err = tx.QueryRow(ctx, ` - INSERT INTO models (public_id, owned_by, input_price_micros_per_million, output_price_micros_per_million, - cache_read_price_micros_per_million, cache_write_price_micros_per_million) - VALUES ($1, $2, $3, $4, $5, $6) - RETURNING id::text, public_id, owned_by, input_price_micros_per_million, output_price_micros_per_million, + INSERT INTO models (public_id, display_name, description, owned_by, input_modalities, output_modalities, + context_window, max_output_tokens, capabilities, regions, lifecycle, released_at, + deprecated_at, retired_at, replacement_model, input_price_micros_per_million, + output_price_micros_per_million, cache_read_price_micros_per_million, cache_write_price_micros_per_million) + VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,NULLIF($15,''),$16,$17,$18,$19) + RETURNING id::text, public_id, display_name, description, owned_by, + input_price_micros_per_million, output_price_micros_per_million, cache_read_price_micros_per_million, cache_write_price_micros_per_million, enabled, created_at`, - input.PublicID, input.OwnedBy, input.InputPriceMicrosPerMillion, input.OutputPriceMicrosPerMillion, + input.PublicID, input.DisplayName, input.Description, input.OwnedBy, inputModalitiesJSON, outputModalitiesJSON, + input.ContextWindow, input.MaxOutputTokens, capabilitiesJSON, regionsJSON, input.Lifecycle, + input.ReleasedAt, input.DeprecatedAt, input.RetiredAt, input.ReplacementModel, + input.InputPriceMicrosPerMillion, input.OutputPriceMicrosPerMillion, input.CacheReadPriceMicrosPerMillion, input.CacheWritePriceMicrosPerMillion, - ).Scan(&result.ID, &result.PublicID, &result.OwnedBy, &result.InputPriceMicrosPerMillion, + ).Scan(&result.ID, &result.PublicID, &result.DisplayName, &result.Description, &result.OwnedBy, &result.InputPriceMicrosPerMillion, &result.OutputPriceMicrosPerMillion, &result.CacheReadPriceMicrosPerMillion, &result.CacheWritePriceMicrosPerMillion, &result.Enabled, &result.CreatedAt) if err != nil { return Model{}, 0, fmt.Errorf("create model: %w", err) } + result.InputModalities, result.OutputModalities = input.InputModalities, input.OutputModalities + result.ContextWindow, result.MaxOutputTokens = input.ContextWindow, input.MaxOutputTokens + result.Capabilities, result.Regions, result.Lifecycle = input.Capabilities, input.Regions, input.Lifecycle + result.ReleasedAt, result.DeprecatedAt, result.RetiredAt = input.ReleasedAt, input.DeprecatedAt, input.RetiredAt + result.ReplacementModel, result.Aliases = input.ReplacementModel, input.Aliases + result.AllowedTenantIDs, result.AllowedKeyIDs = input.AllowedTenantIDs, input.AllowedKeyIDs + result.PriceCurrency, result.PriceVersion = input.PriceCurrency, 1 + err = tx.QueryRow(ctx, `INSERT INTO model_price_versions (model_id,version,currency, + input_price_micros_per_million,output_price_micros_per_million, + cache_read_price_micros_per_million,cache_write_price_micros_per_million) + VALUES ($1,1,$2,$3,$4,$5,$6) RETURNING id::text,effective_from`, result.ID, input.PriceCurrency, + input.InputPriceMicrosPerMillion, input.OutputPriceMicrosPerMillion, + input.CacheReadPriceMicrosPerMillion, input.CacheWritePriceMicrosPerMillion, + ).Scan(&result.PriceVersionID, &result.PriceEffectiveFrom) + if err != nil { + return Model{}, 0, fmt.Errorf("create model price version: %w", err) + } + for _, alias := range input.Aliases { + if _, err := tx.Exec(ctx, `INSERT INTO model_aliases (alias,model_id,deprecated) VALUES ($1,$2,true)`, alias, result.ID); err != nil { + return Model{}, 0, fmt.Errorf("create model alias: %w", err) + } + } + for _, tenantID := range input.AllowedTenantIDs { + if _, err := tx.Exec(ctx, `INSERT INTO model_tenant_allowlist (model_id,tenant_id) VALUES ($1,$2)`, result.ID, tenantID); err != nil { + return Model{}, 0, fmt.Errorf("create tenant model allowlist: %w", err) + } + } + for _, keyID := range input.AllowedKeyIDs { + if _, err := tx.Exec(ctx, `INSERT INTO api_key_model_allowlist (api_key_id,model_id) VALUES ($1,$2)`, keyID, result.ID); err != nil { + return Model{}, 0, fmt.Errorf("create key model allowlist: %w", err) + } + } result.Routes = make([]Route, 0, len(input.Routes)) for _, route := range input.Routes { var created Route @@ -233,6 +311,59 @@ func (s *Store) CreateModel(ctx context.Context, input CreateModelInput) (Model, return result, generation, nil } +func (s *Store) CreateModelPriceVersion(ctx context.Context, modelID string, input CreatePriceVersionInput) (int64, error) { + input.Currency = strings.ToLower(strings.TrimSpace(input.Currency)) + if len(input.Currency) != 3 || input.InputPriceMicrosPerMillion < 0 || input.OutputPriceMicrosPerMillion < 0 || + input.CacheReadPriceMicrosPerMillion < 0 || input.CacheWritePriceMicrosPerMillion < 0 { + return 0, errors.New("price version has invalid currency or negative price") + } + if input.EffectiveFrom.IsZero() { + input.EffectiveFrom = time.Now().UTC() + } + tx, err := s.db.Begin(ctx) + if err != nil { + return 0, err + } + defer tx.Rollback(ctx) + // Lock the model row before allocating the next version. PostgreSQL does + // not allow FOR UPDATE on an aggregate query. + var modelExists bool + if err := tx.QueryRow(ctx, `SELECT true FROM models WHERE id=$1 FOR UPDATE`, modelID).Scan(&modelExists); errors.Is(err, pgx.ErrNoRows) { + return 0, ErrNotFound + } else if err != nil { + return 0, err + } + var currentEffectiveFrom time.Time + err = tx.QueryRow(ctx, `SELECT effective_from FROM model_price_versions WHERE model_id=$1 AND effective_to IS NULL`, modelID).Scan(¤tEffectiveFrom) + if err != nil && !errors.Is(err, pgx.ErrNoRows) { + return 0, err + } + if err == nil && !input.EffectiveFrom.After(currentEffectiveFrom) { + return 0, errors.New("new price version must become effective after the current open version") + } + var version int + if err := tx.QueryRow(ctx, `SELECT COALESCE(max(version),0)+1 FROM model_price_versions WHERE model_id=$1`, modelID).Scan(&version); err != nil { + return 0, err + } + if _, err := tx.Exec(ctx, `UPDATE model_price_versions SET effective_to=$2 WHERE model_id=$1 AND effective_to IS NULL AND effective_from < $2`, modelID, input.EffectiveFrom); err != nil { + return 0, err + } + if _, err := tx.Exec(ctx, `INSERT INTO model_price_versions (model_id,version,currency,input_price_micros_per_million, + output_price_micros_per_million,cache_read_price_micros_per_million,cache_write_price_micros_per_million,effective_from) + VALUES ($1,$2,$3,$4,$5,$6,$7,$8)`, modelID, version, input.Currency, input.InputPriceMicrosPerMillion, + input.OutputPriceMicrosPerMillion, input.CacheReadPriceMicrosPerMillion, input.CacheWritePriceMicrosPerMillion, input.EffectiveFrom); err != nil { + return 0, err + } + generation, err := bumpGeneration(ctx, tx) + if err != nil { + return 0, err + } + if err := tx.Commit(ctx); err != nil { + return 0, err + } + return generation, nil +} + func (s *Store) SetModelEnabled(ctx context.Context, id string, enabled bool) (int64, error) { return s.toggle(ctx, `UPDATE models SET enabled = $2, updated_at = now() WHERE id = $1`, id, enabled) } diff --git a/internal/controlplane/outbox.go b/internal/controlplane/outbox.go index b54a064..b708a90 100644 --- a/internal/controlplane/outbox.go +++ b/internal/controlplane/outbox.go @@ -109,8 +109,9 @@ func (s *Store) ClaimMail(ctx context.Context) (mailer.Message, bool, error) { var ciphertext []byte err = tx.QueryRow(ctx, `WITH candidate AS ( SELECT id FROM console_mail_outbox - WHERE ((status IN ('pending','failed') AND available_at <= now()) + WHERE ((status IN ('pending','retry') AND available_at <= now()) OR (status='sending' AND claimed_at < now()-interval '5 minutes')) + AND NOT EXISTS (SELECT 1 FROM mail_suppressions s WHERE lower(s.recipient)=lower(console_mail_outbox.recipient)) ORDER BY available_at,created_at FOR UPDATE SKIP LOCKED LIMIT 1 ) UPDATE console_mail_outbox o SET status='sending',claimed_at=now(),attempts=attempts+1,last_error='' FROM candidate WHERE o.id=candidate.id @@ -146,7 +147,19 @@ func (s *Store) MarkMailFailed(ctx context.Context, id string, deliveryErr error if len(message) > 1000 { message = message[:1000] } - _, err := s.db.Exec(ctx, `UPDATE console_mail_outbox SET status='failed',claimed_at=NULL,last_error=$2, - available_at=now()+make_interval(secs => LEAST(300, 5 * attempts)) WHERE id=$1 AND status='sending'`, id, message) + _, err := s.db.Exec(ctx, `UPDATE console_mail_outbox SET + status=CASE WHEN attempts>=10 THEN 'dead' ELSE 'retry' END,claimed_at=NULL,last_error=$2, + available_at=now()+make_interval(secs => LEAST(3600, 5 * power(2,LEAST(attempts,9))::int)) + WHERE id=$1 AND status='sending'`, id, message) return err } + +func (s *Store) MailQueueStatus(ctx context.Context) (MailQueueStatus, error) { + var result MailQueueStatus + err := s.db.QueryRow(ctx, `SELECT + count(*) FILTER (WHERE status IN ('pending','sending','retry')), + count(*) FILTER (WHERE status='dead'), + min(created_at) FILTER (WHERE status IN ('pending','sending','retry')) + FROM console_mail_outbox`).Scan(&result.Backlog, &result.Failed, &result.OldestPending) + return result, err +} diff --git a/internal/controlplane/queries.go b/internal/controlplane/queries.go index 9b76fad..73d869f 100644 --- a/internal/controlplane/queries.go +++ b/internal/controlplane/queries.go @@ -173,10 +173,17 @@ func (s *Store) ListProviders(ctx context.Context) ([]Provider, error) { func (s *Store) ListModels(ctx context.Context) ([]Model, error) { rows, err := s.db.Query(ctx, ` - SELECT id::text, public_id, owned_by, input_price_micros_per_million, - output_price_micros_per_million, cache_read_price_micros_per_million, - cache_write_price_micros_per_million, enabled, created_at - FROM models ORDER BY public_id`) + SELECT m.id::text, m.public_id, m.display_name, m.description, m.owned_by, + m.input_modalities, m.output_modalities, m.context_window, m.max_output_tokens, + m.capabilities, m.regions, m.lifecycle, m.released_at, m.deprecated_at, m.retired_at, + COALESCE(m.replacement_model,''), pv.id::text, pv.version, pv.currency, pv.effective_from, + pv.input_price_micros_per_million, pv.output_price_micros_per_million, + pv.cache_read_price_micros_per_million, pv.cache_write_price_micros_per_million, + m.enabled, m.created_at + FROM models m JOIN LATERAL ( + SELECT * FROM model_price_versions v WHERE v.model_id=m.id + ORDER BY v.effective_from DESC LIMIT 1 + ) pv ON TRUE ORDER BY m.public_id`) if err != nil { return nil, fmt.Errorf("query models: %w", err) } @@ -184,13 +191,23 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) { positions := make(map[string]int) for rows.Next() { var item Model - if err := rows.Scan(&item.ID, &item.PublicID, &item.OwnedBy, &item.InputPriceMicrosPerMillion, + var inputModalitiesJSON, outputModalitiesJSON, capabilitiesJSON, regionsJSON []byte + if err := rows.Scan(&item.ID, &item.PublicID, &item.DisplayName, &item.Description, &item.OwnedBy, + &inputModalitiesJSON, &outputModalitiesJSON, &item.ContextWindow, &item.MaxOutputTokens, + &capabilitiesJSON, ®ionsJSON, &item.Lifecycle, &item.ReleasedAt, &item.DeprecatedAt, + &item.RetiredAt, &item.ReplacementModel, &item.PriceVersionID, &item.PriceVersion, + &item.PriceCurrency, &item.PriceEffectiveFrom, &item.InputPriceMicrosPerMillion, &item.OutputPriceMicrosPerMillion, &item.CacheReadPriceMicrosPerMillion, &item.CacheWritePriceMicrosPerMillion, &item.Enabled, &item.CreatedAt); err != nil { rows.Close() return nil, fmt.Errorf("scan model: %w", err) } + _ = json.Unmarshal(inputModalitiesJSON, &item.InputModalities) + _ = json.Unmarshal(outputModalitiesJSON, &item.OutputModalities) + _ = json.Unmarshal(capabilitiesJSON, &item.Capabilities) + _ = json.Unmarshal(regionsJSON, &item.Regions) item.Routes = []Route{} + item.Aliases, item.AllowedTenantIDs, item.AllowedKeyIDs = []string{}, []string{}, []string{} positions[item.ID] = len(models) models = append(models, item) } @@ -200,6 +217,52 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) { } rows.Close() + aliasRows, err := s.db.Query(ctx, `SELECT model_id::text, alias FROM model_aliases ORDER BY alias`) + if err != nil { + return nil, fmt.Errorf("query model aliases: %w", err) + } + for aliasRows.Next() { + var modelID, alias string + if err := aliasRows.Scan(&modelID, &alias); err != nil { + aliasRows.Close() + return nil, err + } + if p, ok := positions[modelID]; ok { + models[p].Aliases = append(models[p].Aliases, alias) + } + } + aliasRows.Close() + tenantRows, err := s.db.Query(ctx, `SELECT model_id::text, tenant_id::text FROM model_tenant_allowlist`) + if err != nil { + return nil, fmt.Errorf("query tenant model allowlist: %w", err) + } + for tenantRows.Next() { + var modelID, tenantID string + if err := tenantRows.Scan(&modelID, &tenantID); err != nil { + tenantRows.Close() + return nil, err + } + if p, ok := positions[modelID]; ok { + models[p].AllowedTenantIDs = append(models[p].AllowedTenantIDs, tenantID) + } + } + tenantRows.Close() + keyRows, err := s.db.Query(ctx, `SELECT model_id::text, api_key_id::text FROM api_key_model_allowlist`) + if err != nil { + return nil, fmt.Errorf("query key model allowlist: %w", err) + } + for keyRows.Next() { + var modelID, keyID string + if err := keyRows.Scan(&modelID, &keyID); err != nil { + keyRows.Close() + return nil, err + } + if p, ok := positions[modelID]; ok { + models[p].AllowedKeyIDs = append(models[p].AllowedKeyIDs, keyID) + } + } + keyRows.Close() + routeRows, err := s.db.Query(ctx, ` SELECT r.id::text, r.model_id::text, r.provider_id::text, p.name, p.protocol, r.upstream_model, r.priority, r.weight, r.enabled diff --git a/internal/controlplane/retention.go b/internal/controlplane/retention.go new file mode 100644 index 0000000..4b20611 --- /dev/null +++ b/internal/controlplane/retention.go @@ -0,0 +1,57 @@ +package controlplane + +import ( + "context" + "log/slog" + "time" +) + +func (s *Store) RunRetentionWorker(ctx context.Context, auditDays, securityDays int, logger *slog.Logger) { + run := func() { + cleanupCtx, cancel := context.WithTimeout(ctx, 30*time.Second) + defer cancel() + if err := s.pruneExpiredSecurityData(cleanupCtx, auditDays, securityDays); err != nil { + logger.Error("retention_cleanup_failed", "error", err) + } + } + run() + ticker := time.NewTicker(24 * time.Hour) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + run() + } + } +} + +func (s *Store) pruneExpiredSecurityData(ctx context.Context, auditDays, securityDays int) error { + tx, err := s.db.Begin(ctx) + if err != nil { + return err + } + defer tx.Rollback(ctx) + auditCutoff := time.Now().UTC().AddDate(0, 0, -auditDays) + securityCutoff := time.Now().UTC().AddDate(0, 0, -securityDays) + queries := []struct { + query string + cutoff time.Time + }{ + {`DELETE FROM audit_logs WHERE created_at < $1`, auditCutoff}, + {`DELETE FROM console_sessions WHERE expires_at < $1 OR (revoked_at IS NOT NULL AND revoked_at < $1)`, securityCutoff}, + {`DELETE FROM console_action_tokens WHERE expires_at < $1 OR (consumed_at IS NOT NULL AND consumed_at < $1)`, securityCutoff}, + {`DELETE FROM console_auth_challenges WHERE expires_at < $1 OR (consumed_at IS NOT NULL AND consumed_at < $1)`, securityCutoff}, + {`DELETE FROM console_webauthn_challenges WHERE expires_at < $1 OR (consumed_at IS NOT NULL AND consumed_at < $1)`, securityCutoff}, + {`DELETE FROM console_login_throttles WHERE updated_at < $1`, securityCutoff}, + {`DELETE FROM console_rate_limits WHERE updated_at < $1`, securityCutoff}, + {`DELETE FROM console_mail_outbox WHERE status='sent' AND sent_at < $1`, securityCutoff}, + } + for _, item := range queries { + if _, err := tx.Exec(ctx, item.query, item.cutoff); err != nil { + return err + } + } + return tx.Commit(ctx) +} diff --git a/internal/controlplane/rotation.go b/internal/controlplane/rotation.go new file mode 100644 index 0000000..1fc8084 --- /dev/null +++ b/internal/controlplane/rotation.go @@ -0,0 +1,69 @@ +package controlplane + +import ( + "context" + "fmt" +) + +type encryptedColumn struct{ table, key, column string } + +// RotateCredentials re-encrypts every control-plane secret with the primary +// key in the configured keyring. Run all gateway instances with both new and +// previous keys before invoking this operation. +func (s *Store) RotateCredentials(ctx context.Context) (int, error) { + tx, err := s.db.Begin(ctx) + if err != nil { + return 0, err + } + defer tx.Rollback(ctx) + columns := []encryptedColumn{ + {"providers", "id", "api_key_ciphertext"}, + {"console_mail_outbox", "id", "body_ciphertext"}, + {"console_totp_credentials", "user_id", "secret_ciphertext"}, + {"console_passkeys", "id", "credential_ciphertext"}, + {"console_webauthn_challenges", "id", "session_ciphertext"}, + } + total := 0 + for _, item := range columns { + rows, queryErr := tx.Query(ctx, fmt.Sprintf(`SELECT %s::text,%s FROM %s`, item.key, item.column, item.table)) + if queryErr != nil { + return total, queryErr + } + type record struct { + id string + ciphertext []byte + } + records := []record{} + for rows.Next() { + var value record + if err := rows.Scan(&value.id, &value.ciphertext); err != nil { + rows.Close() + return total, err + } + records = append(records, value) + } + if err := rows.Err(); err != nil { + rows.Close() + return total, err + } + rows.Close() + for _, value := range records { + plaintext, err := s.cipher.Decrypt(value.ciphertext) + if err != nil { + return total, fmt.Errorf("decrypt %s %s: %w", item.table, value.id, err) + } + ciphertext, err := s.cipher.Encrypt(plaintext) + if err != nil { + return total, err + } + if _, err := tx.Exec(ctx, fmt.Sprintf(`UPDATE %s SET %s=$2 WHERE %s=$1`, item.table, item.column, item.key), value.id, ciphertext); err != nil { + return total, err + } + total++ + } + } + if err := tx.Commit(ctx); err != nil { + return total, err + } + return total, nil +} diff --git a/internal/controlplane/schema.sql b/internal/controlplane/schema.sql index 8a5a605..88db49b 100644 --- a/internal/controlplane/schema.sql +++ b/internal/controlplane/schema.sql @@ -1,5 +1,12 @@ CREATE EXTENSION IF NOT EXISTS pgcrypto; +CREATE TABLE IF NOT EXISTS schema_migrations ( + version BIGINT PRIMARY KEY, + name TEXT NOT NULL, + checksum TEXT NOT NULL, + applied_at TIMESTAMPTZ NOT NULL DEFAULT now() +); + CREATE TABLE IF NOT EXISTS control_state ( singleton BOOLEAN PRIMARY KEY DEFAULT TRUE CHECK (singleton), generation BIGINT NOT NULL DEFAULT 0, @@ -71,6 +78,72 @@ ALTER TABLE models ADD COLUMN IF NOT EXISTS input_price_micros_per_million BIGIN ALTER TABLE models ADD COLUMN IF NOT EXISTS output_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (output_price_micros_per_million >= 0); ALTER TABLE models ADD COLUMN IF NOT EXISTS cache_read_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (cache_read_price_micros_per_million >= 0); ALTER TABLE models ADD COLUMN IF NOT EXISTS cache_write_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (cache_write_price_micros_per_million >= 0); +ALTER TABLE models ADD COLUMN IF NOT EXISTS display_name TEXT NOT NULL DEFAULT ''; +ALTER TABLE models ADD COLUMN IF NOT EXISTS description TEXT NOT NULL DEFAULT ''; +ALTER TABLE models ADD COLUMN IF NOT EXISTS input_modalities JSONB NOT NULL DEFAULT '["text"]'::jsonb; +ALTER TABLE models ADD COLUMN IF NOT EXISTS output_modalities JSONB NOT NULL DEFAULT '["text"]'::jsonb; +ALTER TABLE models ADD COLUMN IF NOT EXISTS context_window BIGINT NOT NULL DEFAULT 0 CHECK (context_window >= 0); +ALTER TABLE models ADD COLUMN IF NOT EXISTS max_output_tokens BIGINT NOT NULL DEFAULT 0 CHECK (max_output_tokens >= 0); +ALTER TABLE models ADD COLUMN IF NOT EXISTS capabilities JSONB NOT NULL DEFAULT '["chat","streaming"]'::jsonb; +ALTER TABLE models ADD COLUMN IF NOT EXISTS regions JSONB NOT NULL DEFAULT '[]'::jsonb; +ALTER TABLE models ADD COLUMN IF NOT EXISTS lifecycle TEXT NOT NULL DEFAULT 'active'; +ALTER TABLE models ADD COLUMN IF NOT EXISTS released_at TIMESTAMPTZ; +ALTER TABLE models ADD COLUMN IF NOT EXISTS deprecated_at TIMESTAMPTZ; +ALTER TABLE models ADD COLUMN IF NOT EXISTS retired_at TIMESTAMPTZ; +ALTER TABLE models ADD COLUMN IF NOT EXISTS replacement_model TEXT; +ALTER TABLE models DROP CONSTRAINT IF EXISTS models_lifecycle_check; +ALTER TABLE models ADD CONSTRAINT models_lifecycle_check CHECK (lifecycle IN ('preview','active','deprecated','retired')); + +CREATE TABLE IF NOT EXISTS model_price_versions ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE, + version INTEGER NOT NULL CHECK (version > 0), + currency TEXT NOT NULL CHECK (currency = lower(currency) AND length(currency) = 3), + input_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (input_price_micros_per_million >= 0), + output_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (output_price_micros_per_million >= 0), + cache_read_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (cache_read_price_micros_per_million >= 0), + cache_write_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (cache_write_price_micros_per_million >= 0), + effective_from TIMESTAMPTZ NOT NULL DEFAULT now(), + effective_to TIMESTAMPTZ, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + UNIQUE (model_id, version), + CHECK (effective_to IS NULL OR effective_to > effective_from) +); +CREATE UNIQUE INDEX IF NOT EXISTS model_price_versions_one_open_idx + ON model_price_versions (model_id) WHERE effective_to IS NULL; +CREATE INDEX IF NOT EXISTS model_price_versions_effective_idx + ON model_price_versions (model_id, effective_from DESC); + +INSERT INTO model_price_versions ( + model_id, version, currency, input_price_micros_per_million, + output_price_micros_per_million, cache_read_price_micros_per_million, + cache_write_price_micros_per_million, effective_from) +SELECT id, 1, 'usd', input_price_micros_per_million, output_price_micros_per_million, + cache_read_price_micros_per_million, cache_write_price_micros_per_million, created_at +FROM models +ON CONFLICT (model_id, version) DO NOTHING; + +CREATE TABLE IF NOT EXISTS model_aliases ( + alias TEXT PRIMARY KEY, + model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE, + deprecated BOOLEAN NOT NULL DEFAULT FALSE, + created_at TIMESTAMPTZ NOT NULL DEFAULT now() +); +CREATE INDEX IF NOT EXISTS model_aliases_model_idx ON model_aliases (model_id); + +CREATE TABLE IF NOT EXISTS model_tenant_allowlist ( + model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE, + tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE CASCADE, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + PRIMARY KEY (model_id, tenant_id) +); + +CREATE TABLE IF NOT EXISTS api_key_model_allowlist ( + api_key_id UUID NOT NULL REFERENCES api_keys(id) ON DELETE CASCADE, + model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + PRIMARY KEY (api_key_id, model_id) +); CREATE TABLE IF NOT EXISTS model_routes ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), @@ -106,7 +179,7 @@ CREATE TABLE IF NOT EXISTS billing_reservations ( public_model TEXT NOT NULL, currency TEXT NOT NULL, reserved_micros BIGINT NOT NULL CHECK (reserved_micros >= 0), - status TEXT NOT NULL DEFAULT 'pending' CHECK (status IN ('pending', 'settled', 'released')), + status TEXT NOT NULL DEFAULT 'pending' CHECK (status IN ('pending', 'settled', 'released', 'metering_failed')), input_price_micros_per_million BIGINT NOT NULL, output_price_micros_per_million BIGINT NOT NULL, cache_read_price_micros_per_million BIGINT NOT NULL, @@ -118,6 +191,30 @@ CREATE TABLE IF NOT EXISTS billing_reservations ( settled_at TIMESTAMPTZ, FOREIGN KEY (project_id, tenant_id) REFERENCES projects(id, tenant_id) ON DELETE CASCADE ); +ALTER TABLE billing_reservations ADD COLUMN IF NOT EXISTS price_version_id UUID REFERENCES model_price_versions(id) ON DELETE SET NULL; +ALTER TABLE billing_reservations DROP CONSTRAINT IF EXISTS billing_reservations_status_check; +ALTER TABLE billing_reservations ADD CONSTRAINT billing_reservations_status_check + CHECK (status IN ('pending','settled','released','metering_failed')); + +CREATE TABLE IF NOT EXISTS billing_settlement_jobs ( + request_id TEXT PRIMARY KEY REFERENCES billing_reservations(request_id) ON DELETE CASCADE, + event JSONB, + status TEXT NOT NULL DEFAULT 'awaiting_event' + CHECK (status IN ('awaiting_event', 'pending', 'processing', 'retry', 'done')), + attempts INTEGER NOT NULL DEFAULT 0 CHECK (attempts >= 0), + available_at TIMESTAMPTZ NOT NULL DEFAULT now(), + locked_at TIMESTAMPTZ, + last_error TEXT NOT NULL DEFAULT '', + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + completed_at TIMESTAMPTZ +); +CREATE INDEX IF NOT EXISTS billing_settlement_jobs_ready_idx + ON billing_settlement_jobs (available_at, created_at) + WHERE status IN ('pending', 'retry'); +CREATE INDEX IF NOT EXISTS billing_settlement_jobs_stale_idx + ON billing_settlement_jobs (created_at) + WHERE status IN ('awaiting_event', 'processing', 'retry'); CREATE TABLE IF NOT EXISTS usage_events ( request_id TEXT PRIMARY KEY, @@ -143,9 +240,17 @@ CREATE TABLE IF NOT EXISTS usage_events ( cost_micros BIGINT NOT NULL DEFAULT 0, charged_micros BIGINT NOT NULL DEFAULT 0, uncollected_micros BIGINT NOT NULL DEFAULT 0, + usage_reported BOOLEAN NOT NULL DEFAULT FALSE, + metering_status TEXT NOT NULL DEFAULT 'not_billable' + CHECK (metering_status IN ('not_billable','reported','missing','upstream_failed')), created_at TIMESTAMPTZ NOT NULL DEFAULT now(), FOREIGN KEY (project_id, tenant_id) REFERENCES projects(id, tenant_id) ON DELETE RESTRICT ); +ALTER TABLE usage_events ADD COLUMN IF NOT EXISTS usage_reported BOOLEAN NOT NULL DEFAULT FALSE; +ALTER TABLE usage_events ADD COLUMN IF NOT EXISTS metering_status TEXT NOT NULL DEFAULT 'not_billable'; +ALTER TABLE usage_events DROP CONSTRAINT IF EXISTS usage_events_metering_status_check; +ALTER TABLE usage_events ADD CONSTRAINT usage_events_metering_status_check + CHECK (metering_status IN ('not_billable','reported','missing','upstream_failed')); -- Usage persistence is independent from billing. Older installations created this -- foreign key, which prevented recording requests when prepaid billing was disabled. @@ -165,6 +270,9 @@ CREATE TABLE IF NOT EXISTS billing_ledger ( created_at TIMESTAMPTZ NOT NULL DEFAULT now(), UNIQUE (source_type, source_id) ); +ALTER TABLE billing_ledger DROP CONSTRAINT IF EXISTS billing_ledger_kind_check; +ALTER TABLE billing_ledger ADD CONSTRAINT billing_ledger_kind_check + CHECK (kind IN ('topup','usage','adjustment','refund','release','dispute','dispute_reversal')); CREATE TABLE IF NOT EXISTS topup_orders ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), @@ -178,13 +286,127 @@ CREATE TABLE IF NOT EXISTS topup_orders ( created_at TIMESTAMPTZ NOT NULL DEFAULT now(), paid_at TIMESTAMPTZ ); +ALTER TABLE topup_orders DROP CONSTRAINT IF EXISTS topup_orders_status_check; +ALTER TABLE topup_orders ADD CONSTRAINT topup_orders_status_check + CHECK (status IN ('pending','paid','failed','expired','partially_refunded','refunded','disputed','reversed')); +ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS stripe_customer_id TEXT; +ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS stripe_payment_intent_id TEXT; +ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS stripe_charge_id TEXT; +ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS stripe_invoice_id TEXT; +ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS invoice_url TEXT; +ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS invoice_pdf_url TEXT; +ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS receipt_url TEXT; +ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS refunded_micros BIGINT NOT NULL DEFAULT 0 CHECK (refunded_micros >= 0); +ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS disputed_micros BIGINT NOT NULL DEFAULT 0 CHECK (disputed_micros >= 0); +ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS reconciliation_status TEXT NOT NULL DEFAULT 'unknown'; +ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS reconciled_at TIMESTAMPTZ; +ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS reconciliation_error TEXT NOT NULL DEFAULT ''; +ALTER TABLE topup_orders DROP CONSTRAINT IF EXISTS topup_orders_reconciliation_status_check; +ALTER TABLE topup_orders ADD CONSTRAINT topup_orders_reconciliation_status_check + CHECK (reconciliation_status IN ('unknown','ok','repaired','missing','mismatch','resolved')); +CREATE INDEX IF NOT EXISTS topup_orders_payment_intent_idx ON topup_orders (stripe_payment_intent_id) WHERE stripe_payment_intent_id IS NOT NULL; +CREATE INDEX IF NOT EXISTS topup_orders_customer_idx ON topup_orders (stripe_customer_id) WHERE stripe_customer_id IS NOT NULL; + +CREATE TABLE IF NOT EXISTS billing_reconciliation_resolutions ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + topup_order_id UUID NOT NULL UNIQUE REFERENCES topup_orders(id) ON DELETE RESTRICT, + tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE RESTRICT, + actor_id TEXT NOT NULL DEFAULT '', + actor_type TEXT NOT NULL CHECK (actor_type IN ('console_user','bootstrap','maintenance')), + reason TEXT NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT now() +); + +CREATE TABLE IF NOT EXISTS stripe_customers ( + tenant_id UUID PRIMARY KEY REFERENCES tenants(id) ON DELETE CASCADE, + stripe_customer_id TEXT NOT NULL UNIQUE, + email TEXT NOT NULL DEFAULT '', + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now() +); + +CREATE TABLE IF NOT EXISTS stripe_refunds ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE RESTRICT, + topup_order_id UUID NOT NULL REFERENCES topup_orders(id) ON DELETE RESTRICT, + stripe_refund_id TEXT UNIQUE, + amount_minor BIGINT NOT NULL CHECK (amount_minor > 0), + amount_micros BIGINT NOT NULL CHECK (amount_micros > 0), + held_micros BIGINT NOT NULL DEFAULT 0 CHECK (held_micros >= 0), + uncollected_micros BIGINT NOT NULL DEFAULT 0 CHECK (uncollected_micros >= 0), + currency TEXT NOT NULL, + reason TEXT NOT NULL DEFAULT 'requested_by_customer', + status TEXT NOT NULL DEFAULT 'queued' + CHECK (status IN ('queued','submitting','pending','requires_action','succeeded','failed','canceled')), + attempts INTEGER NOT NULL DEFAULT 0, + available_at TIMESTAMPTZ NOT NULL DEFAULT now(), + last_error TEXT NOT NULL DEFAULT '', + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + completed_at TIMESTAMPTZ +); +CREATE INDEX IF NOT EXISTS stripe_refunds_ready_idx ON stripe_refunds (available_at, created_at) + WHERE status IN ('queued','submitting'); + +CREATE TABLE IF NOT EXISTS stripe_disputes ( + stripe_dispute_id TEXT PRIMARY KEY, + tenant_id UUID REFERENCES tenants(id) ON DELETE SET NULL, + topup_order_id UUID REFERENCES topup_orders(id) ON DELETE SET NULL, + stripe_payment_intent_id TEXT, + amount_minor BIGINT NOT NULL CHECK (amount_minor >= 0), + amount_micros BIGINT NOT NULL CHECK (amount_micros >= 0), + currency TEXT NOT NULL, + status TEXT NOT NULL, + reason TEXT NOT NULL DEFAULT '', + debited_micros BIGINT NOT NULL DEFAULT 0 CHECK (debited_micros >= 0), + uncollected_micros BIGINT NOT NULL DEFAULT 0 CHECK (uncollected_micros >= 0), + due_by TIMESTAMPTZ, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + closed_at TIMESTAMPTZ +); + +CREATE TABLE IF NOT EXISTS stripe_invoices ( + stripe_invoice_id TEXT PRIMARY KEY, + tenant_id UUID REFERENCES tenants(id) ON DELETE SET NULL, + topup_order_id UUID REFERENCES topup_orders(id) ON DELETE SET NULL, + stripe_customer_id TEXT, + status TEXT NOT NULL DEFAULT '', + currency TEXT NOT NULL DEFAULT '', + amount_due_minor BIGINT NOT NULL DEFAULT 0, + amount_paid_minor BIGINT NOT NULL DEFAULT 0, + attempt_count INTEGER NOT NULL DEFAULT 0, + next_payment_attempt TIMESTAMPTZ, + hosted_invoice_url TEXT NOT NULL DEFAULT '', + invoice_pdf_url TEXT NOT NULL DEFAULT '', + last_failure TEXT NOT NULL DEFAULT '', + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now() +); + +CREATE TABLE IF NOT EXISTS billing_reconciliation_runs ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + status TEXT NOT NULL CHECK (status IN ('running','clean','mismatch','failed')), + checked_orders BIGINT NOT NULL DEFAULT 0, + mismatch_count BIGINT NOT NULL DEFAULT 0, + report JSONB NOT NULL DEFAULT '[]'::jsonb, + error TEXT NOT NULL DEFAULT '', + started_at TIMESTAMPTZ NOT NULL DEFAULT now(), + completed_at TIMESTAMPTZ +); CREATE TABLE IF NOT EXISTS stripe_webhook_events ( event_id TEXT PRIMARY KEY, event_type TEXT NOT NULL, processed_at TIMESTAMPTZ, + attempts INTEGER NOT NULL DEFAULT 0, + last_attempt_at TIMESTAMPTZ, + processing_error TEXT NOT NULL DEFAULT '', created_at TIMESTAMPTZ NOT NULL DEFAULT now() ); +ALTER TABLE stripe_webhook_events ADD COLUMN IF NOT EXISTS attempts INTEGER NOT NULL DEFAULT 0; +ALTER TABLE stripe_webhook_events ADD COLUMN IF NOT EXISTS last_attempt_at TIMESTAMPTZ; +ALTER TABLE stripe_webhook_events ADD COLUMN IF NOT EXISTS processing_error TEXT NOT NULL DEFAULT ''; CREATE TABLE IF NOT EXISTS console_users ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), @@ -279,7 +501,7 @@ CREATE INDEX IF NOT EXISTS console_action_tokens_active_idx CREATE TABLE IF NOT EXISTS console_mail_outbox ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), recipient TEXT NOT NULL, - template TEXT NOT NULL CHECK (template IN ('verify_email', 'password_reset', 'invite')), + template TEXT NOT NULL CHECK (template IN ('verify_email', 'password_reset', 'invite', 'low_balance', 'spend_anomaly')), subject TEXT NOT NULL, body_ciphertext BYTEA NOT NULL, status TEXT NOT NULL DEFAULT 'pending' CHECK (status IN ('pending', 'sending', 'sent', 'failed')), @@ -290,8 +512,43 @@ CREATE TABLE IF NOT EXISTS console_mail_outbox ( last_error TEXT NOT NULL DEFAULT '', created_at TIMESTAMPTZ NOT NULL DEFAULT now() ); -CREATE INDEX IF NOT EXISTS console_mail_outbox_pending_idx - ON console_mail_outbox (available_at, created_at) WHERE status IN ('pending', 'failed'); +ALTER TABLE console_mail_outbox DROP CONSTRAINT IF EXISTS console_mail_outbox_template_check; +ALTER TABLE console_mail_outbox ADD CONSTRAINT console_mail_outbox_template_check + CHECK (template IN ('verify_email','password_reset','invite','low_balance','spend_anomaly')); +ALTER TABLE console_mail_outbox DROP CONSTRAINT IF EXISTS console_mail_outbox_status_check; +UPDATE console_mail_outbox SET status='retry' WHERE status='failed'; +ALTER TABLE console_mail_outbox ADD CONSTRAINT console_mail_outbox_status_check + CHECK (status IN ('pending','sending','sent','retry','dead','suppressed')); +DROP INDEX IF EXISTS console_mail_outbox_pending_idx; +CREATE INDEX console_mail_outbox_pending_idx + ON console_mail_outbox (available_at, created_at) WHERE status IN ('pending', 'retry'); + +CREATE TABLE IF NOT EXISTS mail_suppressions ( + recipient TEXT PRIMARY KEY, + reason TEXT NOT NULL CHECK (reason IN ('bounce','complaint','manual')), + provider TEXT NOT NULL DEFAULT '', + provider_event_id TEXT NOT NULL DEFAULT '', + detail TEXT NOT NULL DEFAULT '', + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now() +); + +CREATE TABLE IF NOT EXISTS mail_feedback_events ( + event_id TEXT PRIMARY KEY, + event_type TEXT NOT NULL CHECK (event_type IN ('delivered','bounce','complaint')), + recipient TEXT NOT NULL, + provider TEXT NOT NULL DEFAULT '', + created_at TIMESTAMPTZ NOT NULL DEFAULT now() +); + +CREATE TABLE IF NOT EXISTS mail_notification_events ( + tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE CASCADE, + recipient TEXT NOT NULL, + notification_type TEXT NOT NULL CHECK (notification_type IN ('low_balance','spend_anomaly')), + dedupe_key TEXT NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + PRIMARY KEY (tenant_id,recipient,notification_type,dedupe_key) +); CREATE TABLE IF NOT EXISTS console_auth_challenges ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), diff --git a/internal/controlplane/snapshot.go b/internal/controlplane/snapshot.go index bc04e7c..b9bddaf 100644 --- a/internal/controlplane/snapshot.go +++ b/internal/controlplane/snapshot.go @@ -6,6 +6,7 @@ import ( "encoding/json" "fmt" "strings" + "time" "aigw/internal/auth" "aigw/internal/domain" @@ -102,13 +103,23 @@ func (s *Store) loadProviders(ctx context.Context, tx pgx.Tx) (map[string]domain func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Provider) ([]domain.Model, error) { rows, err := tx.Query(ctx, ` - SELECT m.public_id, m.owned_by, m.input_price_micros_per_million, m.output_price_micros_per_million, - m.cache_read_price_micros_per_million, m.cache_write_price_micros_per_million, + SELECT m.id::text, m.public_id, m.display_name, m.description, m.owned_by, + m.input_modalities, m.output_modalities, m.context_window, m.max_output_tokens, + m.capabilities, m.regions, m.lifecycle, m.released_at, m.deprecated_at, + m.retired_at, COALESCE(m.replacement_model,''), + pv.id::text, pv.version, pv.currency, pv.effective_from, + pv.input_price_micros_per_million, pv.output_price_micros_per_million, + pv.cache_read_price_micros_per_million, pv.cache_write_price_micros_per_million, r.provider_id::text, r.upstream_model, r.priority, r.weight FROM models m + JOIN LATERAL ( + SELECT * FROM model_price_versions v WHERE v.model_id=m.id + AND v.effective_from <= now() AND (v.effective_to IS NULL OR v.effective_to > now()) + ORDER BY v.effective_from DESC LIMIT 1 + ) pv ON TRUE JOIN model_routes r ON r.model_id = m.id AND r.enabled = TRUE JOIN providers p ON p.id = r.provider_id AND p.enabled = TRUE - WHERE m.enabled = TRUE + WHERE m.enabled = TRUE AND m.lifecycle <> 'retired' ORDER BY m.public_id, r.priority, r.created_at`) if err != nil { return nil, fmt.Errorf("query model routes: %w", err) @@ -117,10 +128,21 @@ func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Prov models := make([]domain.Model, 0) index := make(map[string]int) for rows.Next() { - var publicID, ownedBy, providerID, upstreamModel string + var modelID, publicID, displayName, description, ownedBy, replacement, providerID, upstreamModel string + var inputModalitiesJSON, outputModalitiesJSON, capabilitiesJSON, regionsJSON []byte + var contextWindow, maxOutput int64 + var lifecycle, priceVersionID, priceCurrency string + var releasedAt, deprecatedAt, retiredAt *time.Time + var priceVersion int + var priceEffectiveFrom time.Time var inputPrice, outputPrice, cacheReadPrice, cacheWritePrice int64 var priority, weight int - if err := rows.Scan(&publicID, &ownedBy, &inputPrice, &outputPrice, &cacheReadPrice, &cacheWritePrice, &providerID, &upstreamModel, &priority, &weight); err != nil { + if err := rows.Scan(&modelID, &publicID, &displayName, &description, &ownedBy, + &inputModalitiesJSON, &outputModalitiesJSON, &contextWindow, &maxOutput, + &capabilitiesJSON, ®ionsJSON, &lifecycle, &releasedAt, &deprecatedAt, &retiredAt, &replacement, + &priceVersionID, &priceVersion, &priceCurrency, &priceEffectiveFrom, + &inputPrice, &outputPrice, &cacheReadPrice, &cacheWritePrice, + &providerID, &upstreamModel, &priority, &weight); err != nil { return nil, fmt.Errorf("scan model route: %w", err) } provider, ok := providers[providerID] @@ -129,13 +151,32 @@ func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Prov } position, exists := index[publicID] if !exists { + var inputModalities, outputModalities, capabilities, regions []string + if err := json.Unmarshal(inputModalitiesJSON, &inputModalities); err != nil { + return nil, fmt.Errorf("decode model input modalities: %w", err) + } + if err := json.Unmarshal(outputModalitiesJSON, &outputModalities); err != nil { + return nil, fmt.Errorf("decode model output modalities: %w", err) + } + if err := json.Unmarshal(capabilitiesJSON, &capabilities); err != nil { + return nil, fmt.Errorf("decode model capabilities: %w", err) + } + if err := json.Unmarshal(regionsJSON, ®ions); err != nil { + return nil, fmt.Errorf("decode model regions: %w", err) + } position = len(models) index[publicID] = position models = append(models, domain.Model{ - ID: publicID, OwnedBy: ownedBy, + ID: publicID, DisplayName: displayName, Description: description, OwnedBy: ownedBy, + InputModalities: inputModalities, OutputModalities: outputModalities, + ContextWindow: contextWindow, MaxOutputTokens: maxOutput, Capabilities: capabilities, + Regions: regions, Lifecycle: lifecycle, ReleasedAt: releasedAt, DeprecatedAt: deprecatedAt, + RetiredAt: retiredAt, ReplacementModel: replacement, PriceVersionID: priceVersionID, + PriceVersion: priceVersion, PriceCurrency: priceCurrency, PriceEffectiveFrom: priceEffectiveFrom, InputPriceMicrosPerMillion: inputPrice, OutputPriceMicrosPerMillion: outputPrice, CacheReadPriceMicrosPerMillion: cacheReadPrice, CacheWritePriceMicrosPerMillion: cacheWritePrice, }) + _ = modelID } models[position].Routes = append(models[position].Routes, domain.Route{ Provider: provider, UpstreamModel: upstreamModel, Priority: priority, Weight: weight, @@ -144,6 +185,61 @@ func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Prov if err := rows.Err(); err != nil { return nil, fmt.Errorf("read model routes: %w", err) } + byID := make(map[string]int, len(models)) + for i := range models { + byID[models[i].ID] = i + } + aliasRows, err := tx.Query(ctx, `SELECT a.alias, m.public_id FROM model_aliases a JOIN models m ON m.id=a.model_id`) + if err != nil { + return nil, fmt.Errorf("query model aliases: %w", err) + } + for aliasRows.Next() { + var alias, publicID string + if err := aliasRows.Scan(&alias, &publicID); err != nil { + aliasRows.Close() + return nil, err + } + if position, ok := byID[publicID]; ok { + models[position].Aliases = append(models[position].Aliases, alias) + } + } + aliasRows.Close() + tenantRows, err := tx.Query(ctx, `SELECT m.public_id, a.tenant_id::text FROM model_tenant_allowlist a JOIN models m ON m.id=a.model_id`) + if err != nil { + return nil, fmt.Errorf("query model tenant allowlist: %w", err) + } + for tenantRows.Next() { + var publicID, tenantID string + if err := tenantRows.Scan(&publicID, &tenantID); err != nil { + tenantRows.Close() + return nil, err + } + if position, ok := byID[publicID]; ok { + if models[position].AllowedTenantIDs == nil { + models[position].AllowedTenantIDs = map[string]struct{}{} + } + models[position].AllowedTenantIDs[tenantID] = struct{}{} + } + } + tenantRows.Close() + keyRows, err := tx.Query(ctx, `SELECT m.public_id, a.api_key_id::text FROM api_key_model_allowlist a JOIN models m ON m.id=a.model_id`) + if err != nil { + return nil, fmt.Errorf("query model key allowlist: %w", err) + } + for keyRows.Next() { + var publicID, keyID string + if err := keyRows.Scan(&publicID, &keyID); err != nil { + keyRows.Close() + return nil, err + } + if position, ok := byID[publicID]; ok { + if models[position].AllowedKeyIDs == nil { + models[position].AllowedKeyIDs = map[string]struct{}{} + } + models[position].AllowedKeyIDs[keyID] = struct{}{} + } + } + keyRows.Close() return models, nil } diff --git a/internal/controlplane/store.go b/internal/controlplane/store.go index 833f846..fb48d78 100644 --- a/internal/controlplane/store.go +++ b/internal/controlplane/store.go @@ -2,7 +2,9 @@ package controlplane import ( "context" + "crypto/sha256" _ "embed" + "encoding/hex" "encoding/json" "errors" "fmt" @@ -11,6 +13,7 @@ import ( "aigw/internal/security" + "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" "github.com/redis/go-redis/v9" ) @@ -20,12 +23,15 @@ var schemaSQL string var ErrRedisDisabled = errors.New("Redis propagation is disabled") +const migrationVersion int64 = 2026080504 + type Options struct { - DatabaseURL string - RedisURL string - CredentialKey string - RedisChannel string - VersionCacheKey string + DatabaseURL string + RedisURL string + CredentialKey string + PreviousCredentialKeys []string + RedisChannel string + VersionCacheKey string } type Store struct { @@ -37,7 +43,7 @@ type Store struct { } func NewStore(ctx context.Context, options Options) (*Store, error) { - cipher, err := security.NewCredentialCipher(options.CredentialKey) + cipher, err := security.NewCredentialKeyring(options.CredentialKey, options.PreviousCredentialKeys) if err != nil { return nil, err } @@ -76,11 +82,10 @@ func (s *Store) RedisEnabled() bool { return s.redis != nil } +func (s *Store) Ping(ctx context.Context) error { return s.db.Ping(ctx) } + func (s *Store) Migrate(ctx context.Context) error { - if _, err := s.db.Exec(ctx, schemaSQL); err != nil { - return fmt.Errorf("apply control-plane schema: %w", err) - } - return nil + return applySchema(ctx, s.db) } func MigrateDatabase(ctx context.Context, databaseURL string) error { @@ -89,12 +94,67 @@ func MigrateDatabase(ctx context.Context, databaseURL string) error { return fmt.Errorf("configure PostgreSQL: %w", err) } defer db.Close() - if _, err := db.Exec(ctx, schemaSQL); err != nil { + return applySchema(ctx, db) +} + +func MigrationStatusDatabase(ctx context.Context, databaseURL string) (MigrationStatus, error) { + db, err := pgxpool.New(ctx, databaseURL) + if err != nil { + return MigrationStatus{}, fmt.Errorf("configure PostgreSQL: %w", err) + } + defer db.Close() + var result MigrationStatus + err = db.QueryRow(ctx, `SELECT version,name,checksum,applied_at FROM schema_migrations ORDER BY version DESC LIMIT 1`).Scan(&result.Version, &result.Name, &result.Checksum, &result.AppliedAt) + if err != nil { + return MigrationStatus{}, fmt.Errorf("read migration status: %w", err) + } + return result, nil +} + +func applySchema(ctx context.Context, db *pgxpool.Pool) error { + tx, err := db.Begin(ctx) + if err != nil { + return fmt.Errorf("begin migration: %w", err) + } + defer tx.Rollback(ctx) + if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock($1)`, migrationVersion); err != nil { + return err + } + if _, err := tx.Exec(ctx, schemaSQL); err != nil { return fmt.Errorf("apply control-plane schema: %w", err) } + hash := sha256.Sum256([]byte(schemaSQL)) + checksum := hex.EncodeToString(hash[:]) + var existing string + err = tx.QueryRow(ctx, `SELECT checksum FROM schema_migrations WHERE version=$1`, migrationVersion).Scan(&existing) + if err == nil && existing != checksum { + return fmt.Errorf("migration %d checksum changed; deploy an explicit new migration version", migrationVersion) + } + if !errors.Is(err, pgx.ErrNoRows) && err != nil { + return err + } + if _, err := tx.Exec(ctx, `INSERT INTO schema_migrations(version,name,checksum) VALUES ($1,$2,$3) ON CONFLICT DO NOTHING`, migrationVersion, "commercial-control-plane", checksum); err != nil { + return err + } + if err := tx.Commit(ctx); err != nil { + return fmt.Errorf("commit migration: %w", err) + } return nil } +type MigrationStatus struct { + Version int64 `json:"version"` + Name string `json:"name"` + Checksum string `json:"checksum"` + AppliedAt time.Time `json:"applied_at"` +} + +func (s *Store) MigrationStatus(ctx context.Context) (MigrationStatus, error) { + var result MigrationStatus + err := s.db.QueryRow(ctx, `SELECT version,name,checksum,applied_at FROM schema_migrations ORDER BY version DESC LIMIT 1`).Scan(&result.Version, &result.Name, &result.Checksum, &result.AppliedAt) + return result, err +} + func (s *Store) DatabaseGeneration(ctx context.Context) (int64, error) { var generation int64 err := s.db.QueryRow(ctx, `SELECT generation FROM control_state WHERE singleton = TRUE`).Scan(&generation) diff --git a/internal/controlplane/store_integration_test.go b/internal/controlplane/store_integration_test.go new file mode 100644 index 0000000..4f55e10 --- /dev/null +++ b/internal/controlplane/store_integration_test.go @@ -0,0 +1,66 @@ +package controlplane + +import ( + "context" + "fmt" + "net/url" + "os" + "testing" + "time" + + "github.com/jackc/pgx/v5/pgxpool" +) + +func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) { + databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL") + if databaseURL == "" { + t.Skip("AIGW_TEST_DATABASE_URL is not set") + } + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + db, err := pgxpool.New(ctx, databaseURL) + if err != nil { + t.Fatal(err) + } + defer db.Close() + schema := fmt.Sprintf("migration_drill_%d", time.Now().UnixNano()) + if _, err := db.Exec(ctx, "CREATE SCHEMA "+schema); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _, _ = db.Exec(context.Background(), "DROP SCHEMA "+schema+" CASCADE") }) + if _, err := db.Exec(ctx, `CREATE TABLE `+schema+`.schema_migrations ( + version BIGINT PRIMARY KEY,name TEXT NOT NULL,checksum TEXT NOT NULL,applied_at TIMESTAMPTZ NOT NULL DEFAULT now())`); err != nil { + t.Fatal(err) + } + if _, err := db.Exec(ctx, `INSERT INTO `+schema+`.schema_migrations(version,name,checksum) VALUES ($1,'previous-release','immutable-previous-checksum')`, migrationVersion-1); err != nil { + t.Fatal(err) + } + parsed, err := url.Parse(databaseURL) + if err != nil { + t.Fatal(err) + } + query := parsed.Query() + query.Set("search_path", schema) + parsed.RawQuery = query.Encode() + isolatedURL := parsed.String() + if err := MigrateDatabase(ctx, isolatedURL); err != nil { + t.Fatal(err) + } + if err := MigrateDatabase(ctx, isolatedURL); err != nil { + t.Fatalf("second migration must be idempotent: %v", err) + } + status, err := MigrationStatusDatabase(ctx, isolatedURL) + if err != nil { + t.Fatal(err) + } + if status.Version != migrationVersion { + t.Fatalf("migration version = %d, want %d", status.Version, migrationVersion) + } + var count int + if err := db.QueryRow(ctx, `SELECT count(*) FROM `+schema+`.schema_migrations`).Scan(&count); err != nil { + t.Fatal(err) + } + if count != 2 { + t.Fatalf("migration history contains %d rows, want previous and current", count) + } +} diff --git a/internal/controlplane/types.go b/internal/controlplane/types.go index f81e63b..c24dde9 100644 --- a/internal/controlplane/types.go +++ b/internal/controlplane/types.go @@ -7,6 +7,12 @@ import ( "aigw/internal/domain" ) +type MailQueueStatus struct { + Backlog int64 `json:"backlog"` + Failed int64 `json:"failed"` + OldestPending *time.Time `json:"oldest_pending,omitempty"` +} + type Tenant struct { ID string `json:"id"` Slug string `json:"slug"` @@ -62,16 +68,36 @@ type Route struct { } type Model struct { - ID string `json:"id"` - PublicID string `json:"public_id"` - OwnedBy string `json:"owned_by"` - InputPriceMicrosPerMillion int64 `json:"input_price_micros_per_million"` - OutputPriceMicrosPerMillion int64 `json:"output_price_micros_per_million"` - CacheReadPriceMicrosPerMillion int64 `json:"cache_read_price_micros_per_million"` - CacheWritePriceMicrosPerMillion int64 `json:"cache_write_price_micros_per_million"` - Enabled bool `json:"enabled"` - Routes []Route `json:"routes"` - CreatedAt time.Time `json:"created_at"` + ID string `json:"id"` + PublicID string `json:"public_id"` + DisplayName string `json:"display_name"` + Description string `json:"description"` + OwnedBy string `json:"owned_by"` + InputModalities []string `json:"input_modalities"` + OutputModalities []string `json:"output_modalities"` + ContextWindow int64 `json:"context_window"` + MaxOutputTokens int64 `json:"max_output_tokens"` + Capabilities []string `json:"capabilities"` + Regions []string `json:"regions"` + Lifecycle string `json:"lifecycle"` + ReleasedAt *time.Time `json:"released_at,omitempty"` + DeprecatedAt *time.Time `json:"deprecated_at,omitempty"` + RetiredAt *time.Time `json:"retired_at,omitempty"` + ReplacementModel string `json:"replacement_model,omitempty"` + Aliases []string `json:"aliases"` + AllowedTenantIDs []string `json:"allowed_tenant_ids"` + AllowedKeyIDs []string `json:"allowed_key_ids"` + PriceVersionID string `json:"price_version_id"` + PriceVersion int `json:"price_version"` + PriceCurrency string `json:"price_currency"` + PriceEffectiveFrom time.Time `json:"price_effective_from"` + InputPriceMicrosPerMillion int64 `json:"input_price_micros_per_million"` + OutputPriceMicrosPerMillion int64 `json:"output_price_micros_per_million"` + CacheReadPriceMicrosPerMillion int64 `json:"cache_read_price_micros_per_million"` + CacheWritePriceMicrosPerMillion int64 `json:"cache_write_price_micros_per_million"` + Enabled bool `json:"enabled"` + Routes []Route `json:"routes"` + CreatedAt time.Time `json:"created_at"` } type Overview struct { @@ -141,7 +167,24 @@ type RouteInput struct { type CreateModelInput struct { PublicID string `json:"public_id"` + DisplayName string `json:"display_name"` + Description string `json:"description"` OwnedBy string `json:"owned_by"` + InputModalities []string `json:"input_modalities"` + OutputModalities []string `json:"output_modalities"` + ContextWindow int64 `json:"context_window"` + MaxOutputTokens int64 `json:"max_output_tokens"` + Capabilities []string `json:"capabilities"` + Regions []string `json:"regions"` + Lifecycle string `json:"lifecycle"` + ReleasedAt *time.Time `json:"released_at"` + DeprecatedAt *time.Time `json:"deprecated_at"` + RetiredAt *time.Time `json:"retired_at"` + ReplacementModel string `json:"replacement_model"` + Aliases []string `json:"aliases"` + AllowedTenantIDs []string `json:"allowed_tenant_ids"` + AllowedKeyIDs []string `json:"allowed_key_ids"` + PriceCurrency string `json:"price_currency"` InputPriceMicrosPerMillion int64 `json:"input_price_micros_per_million"` OutputPriceMicrosPerMillion int64 `json:"output_price_micros_per_million"` CacheReadPriceMicrosPerMillion int64 `json:"cache_read_price_micros_per_million"` @@ -149,6 +192,15 @@ type CreateModelInput struct { Routes []RouteInput `json:"routes"` } +type CreatePriceVersionInput struct { + Currency string `json:"currency"` + InputPriceMicrosPerMillion int64 `json:"input_price_micros_per_million"` + OutputPriceMicrosPerMillion int64 `json:"output_price_micros_per_million"` + CacheReadPriceMicrosPerMillion int64 `json:"cache_read_price_micros_per_million"` + CacheWritePriceMicrosPerMillion int64 `json:"cache_write_price_micros_per_million"` + EffectiveFrom time.Time `json:"effective_from"` +} + type ConsoleActor struct { ID string `json:"id,omitempty"` TenantID string `json:"tenant_id,omitempty"` diff --git a/internal/controlplane/usage.go b/internal/controlplane/usage.go index 6c69a9f..b436d8c 100644 --- a/internal/controlplane/usage.go +++ b/internal/controlplane/usage.go @@ -29,13 +29,13 @@ func (s *Store) RecordUsage(ctx context.Context, event domain.UsageEvent) error request_id, tenant_id, project_id, key_id, public_model, provider_id, upstream_model, protocol, stream, status_code, success, error_type, attempts, started_at, duration_ms, input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens, - cost_micros, charged_micros, uncollected_micros) - VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,0,0,0) + cost_micros, charged_micros, uncollected_micros, usage_reported, metering_status) + VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,0,0,0,$21,$22) ON CONFLICT (request_id) DO NOTHING`, event.RequestID, event.TenantID, event.ProjectID, event.KeyID, event.PublicModel, event.ProviderID, event.UpstreamModel, string(event.Protocol), event.Stream, event.StatusCode, event.Success, event.ErrorType, event.Attempts, event.StartedAt, event.DurationMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens, - event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens) + event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens, event.UsageReported, usageMeteringStatus(event)) if err != nil { return fmt.Errorf("persist usage event: %w", err) } @@ -51,6 +51,16 @@ func (s *Store) RecordUsage(ctx context.Context, event domain.UsageEvent) error return nil } +func usageMeteringStatus(event domain.UsageEvent) string { + if event.StatusCode < 200 || event.StatusCode >= 300 || !event.Success { + return "upstream_failed" + } + if event.UsageReported { + return "reported" + } + return "missing" +} + func upsertUsageRollup(ctx context.Context, tx pgx.Tx, event domain.UsageEvent, cost, charged, uncollected int64) error { period := time.Date(event.StartedAt.UTC().Year(), event.StartedAt.UTC().Month(), 1, 0, 0, 0, 0, time.UTC) _, err := tx.Exec(ctx, ` diff --git a/internal/domain/types.go b/internal/domain/types.go index b1decc9..94e56c1 100644 --- a/internal/domain/types.go +++ b/internal/domain/types.go @@ -32,7 +32,27 @@ type Route struct { type Model struct { ID string + DisplayName string + Description string OwnedBy string + InputModalities []string + OutputModalities []string + ContextWindow int64 + MaxOutputTokens int64 + Capabilities []string + Regions []string + Lifecycle string + ReleasedAt *time.Time + DeprecatedAt *time.Time + RetiredAt *time.Time + ReplacementModel string + Aliases []string + AllowedTenantIDs map[string]struct{} + AllowedKeyIDs map[string]struct{} + PriceVersionID string + PriceVersion int + PriceCurrency string + PriceEffectiveFrom time.Time InputPriceMicrosPerMillion int64 OutputPriceMicrosPerMillion int64 CacheReadPriceMicrosPerMillion int64 @@ -40,6 +60,20 @@ type Model struct { Routes []Route } +func (m Model) Allows(principal Principal) bool { + if len(m.AllowedTenantIDs) > 0 { + if _, ok := m.AllowedTenantIDs[principal.TenantID]; !ok { + return false + } + } + if len(m.AllowedKeyIDs) > 0 { + if _, ok := m.AllowedKeyIDs[principal.KeyID]; !ok { + return false + } + } + return true +} + type Usage struct { InputTokens int64 `json:"input_tokens,omitempty"` OutputTokens int64 `json:"output_tokens,omitempty"` @@ -65,6 +99,7 @@ type UsageEvent struct { StartedAt time.Time `json:"started_at"` DurationMS int64 `json:"duration_ms"` Usage Usage `json:"usage"` + UsageReported bool `json:"usage_reported"` } // LimitPolicy is the immutable runtime view of a project's commercial limits. diff --git a/internal/httpapi/api.go b/internal/httpapi/api.go index becfbf3..6ab437c 100644 --- a/internal/httpapi/api.go +++ b/internal/httpapi/api.go @@ -28,54 +28,57 @@ import ( type requestIDKey struct{} +type UsageRecorder interface { + RecordUsage(context.Context, domain.UsageEvent) error +} + type API struct { - authenticator auth.Authenticator - catalog *catalog.Catalog - router *routing.Router - forwarder *provider.Forwarder - usageSink telemetry.UsageSink - billingMeter billing.Meter - limiter *limits.Limiter - usageRecorder interface { - RecordUsage(context.Context, domain.UsageEvent) error - } - metrics *telemetry.Metrics - logger *slog.Logger - maxBodyBytes int64 - exposeMetrics bool + authenticator auth.Authenticator + catalog *catalog.Catalog + router *routing.Router + forwarder *provider.Forwarder + usageSink telemetry.UsageSink + billingMeter billing.Meter + limiter *limits.Limiter + usageRecorder UsageRecorder + metrics *telemetry.Metrics + logger *slog.Logger + maxBodyBytes int64 + exposeMetrics bool + deploymentRegion string } type Options struct { - Authenticator auth.Authenticator - Catalog *catalog.Catalog - Router *routing.Router - Forwarder *provider.Forwarder - UsageSink telemetry.UsageSink - BillingMeter billing.Meter - Limiter *limits.Limiter - UsageRecorder interface { - RecordUsage(context.Context, domain.UsageEvent) error - } - Metrics *telemetry.Metrics - Logger *slog.Logger - MaxBodyBytes int64 - ExposeMetrics bool + Authenticator auth.Authenticator + Catalog *catalog.Catalog + Router *routing.Router + Forwarder *provider.Forwarder + UsageSink telemetry.UsageSink + BillingMeter billing.Meter + Limiter *limits.Limiter + UsageRecorder UsageRecorder + Metrics *telemetry.Metrics + Logger *slog.Logger + MaxBodyBytes int64 + ExposeMetrics bool + DeploymentRegion string } func New(options Options) *API { return &API{ - authenticator: options.Authenticator, - catalog: options.Catalog, - router: options.Router, - forwarder: options.Forwarder, - usageSink: options.UsageSink, - billingMeter: options.BillingMeter, - limiter: options.Limiter, - usageRecorder: options.UsageRecorder, - metrics: options.Metrics, - logger: options.Logger, - maxBodyBytes: options.MaxBodyBytes, - exposeMetrics: options.ExposeMetrics, + authenticator: options.Authenticator, + catalog: options.Catalog, + router: options.Router, + forwarder: options.Forwarder, + usageSink: options.UsageSink, + billingMeter: options.BillingMeter, + limiter: options.Limiter, + usageRecorder: options.UsageRecorder, + metrics: options.Metrics, + logger: options.Logger, + maxBodyBytes: options.MaxBodyBytes, + exposeMetrics: options.ExposeMetrics, + deploymentRegion: strings.ToLower(strings.TrimSpace(options.DeploymentRegion)), } } @@ -87,6 +90,17 @@ func (a *API) Handler() http.Handler { mux.Handle("GET /metrics", a.metrics) } + a.registerInference(mux) + return a.withRequestID(a.recoverPanics(mux)) +} + +func (a *API) InferenceHandler() http.Handler { + mux := http.NewServeMux() + a.registerInference(mux) + return a.withRequestID(a.recoverPanics(mux)) +} + +func (a *API) registerInference(mux *http.ServeMux) { mux.HandleFunc("GET /v1/models", a.openAIModels) mux.HandleFunc("GET /api/v1/models", a.openAIModels) mux.HandleFunc("POST /v1/chat/completions", a.openAIChat) @@ -97,7 +111,6 @@ func (a *API) Handler() http.Handler { mux.HandleFunc("POST /anthropic/v1/messages", a.anthropicMessages) mux.HandleFunc("POST /api/anthropic/v1/messages", a.anthropicMessages) - return a.withRequestID(a.recoverPanics(mux)) } func (a *API) health(w http.ResponseWriter, _ *http.Request) { @@ -143,8 +156,13 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do } var envelope struct { - Model string `json:"model"` - Stream bool `json:"stream"` + Model string `json:"model"` + Stream bool `json:"stream"` + MaxTokens int64 `json:"max_tokens"` + MaxCompletionTokens int64 `json:"max_completion_tokens"` + MaxOutputTokens int64 `json:"max_output_tokens"` + Tools json.RawMessage `json:"tools"` + ResponseFormat json.RawMessage `json:"response_format"` } if err := json.Unmarshal(body, &envelope); err != nil || strings.TrimSpace(envelope.Model) == "" { apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "invalid_params", Message: "Parameter model is required and the body must be valid JSON"}, requestID) @@ -170,7 +188,30 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do defer lease.Release() } - routes, err := a.router.Plan(envelope.Model, protocol) + model, modelErr := a.catalog.ModelForPrincipal(envelope.Model, principal) + if modelErr != nil || !a.modelAvailableInRegion(model) { + apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "invalid_model", Message: "Model does not exist or is not allowed for this API key"}, requestID) + return + } + requestedOutput := max(envelope.MaxTokens, envelope.MaxCompletionTokens, envelope.MaxOutputTokens) + if model.MaxOutputTokens > 0 && requestedOutput > model.MaxOutputTokens { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "invalid_params", Message: "Requested output exceeds the model maximum"}, requestID) + return + } + if envelope.Stream && !modelHasCapability(model, "streaming") { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "unsupported_capability", Message: "Model does not support streaming"}, requestID) + return + } + if len(envelope.Tools) > 0 && string(envelope.Tools) != "null" && !modelHasCapability(model, "tools") { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "unsupported_capability", Message: "Model does not support tools"}, requestID) + return + } + if len(envelope.ResponseFormat) > 0 && string(envelope.ResponseFormat) != "null" && !modelHasCapability(model, "json") { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "unsupported_capability", Message: "Model does not support structured JSON output"}, requestID) + return + } + + routes, err := a.router.Plan(model.ID, protocol) if err != nil { if errors.Is(err, routing.ErrNoRoute) { apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "model_not_supported", Message: "Model does not support this API protocol"}, requestID) @@ -180,11 +221,6 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do return } if a.billingMeter != nil { - model, modelErr := a.catalog.Model(envelope.Model) - if modelErr != nil { - apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "invalid_model", Message: "Model does not exist"}, requestID) - return - } policy := domain.LimitPolicy{} if a.limiter != nil { policy, _ = a.limiter.Policy(principal.ProjectID) @@ -257,6 +293,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do observer := usage.NewObserver(protocol, stream) copyErr := copyResponse(w, result.Response.Body, observer, stream) usageResult := observer.Usage() + usageReported := observer.Reported() success = copyErr == nil errorType := "" if copyErr != nil { @@ -267,6 +304,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do PublicModel: envelope.Model, ProviderID: result.Route.Provider.ID, UpstreamModel: result.Route.UpstreamModel, Protocol: protocol, Stream: stream, StatusCode: result.Response.StatusCode, Success: success, ErrorType: errorType, Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(), Usage: usageResult, + UsageReported: usageReported, }) a.logger.Info("inference_request", "request_id", requestID, @@ -281,22 +319,27 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do } func (a *API) openAIModels(w http.ResponseWriter, r *http.Request) { - if !a.authorize(w, r) { + principal, ok := a.authorize(w, r) + if !ok { return } - models := a.catalog.Models(domain.ProtocolOpenAI) + models := a.availableModels(a.catalog.ModelsFor(domain.ProtocolOpenAI, principal)) data := make([]map[string]any, 0, len(models)) for _, model := range models { - data = append(data, map[string]any{"id": model.ID, "object": "model", "created": 0, "owned_by": model.OwnedBy}) + data = append(data, map[string]any{"id": model.ID, "object": "model", "created": modelCreated(model), "owned_by": model.OwnedBy, + "display_name": model.DisplayName, "context_window": model.ContextWindow, "max_output_tokens": model.MaxOutputTokens, + "input_modalities": model.InputModalities, "output_modalities": model.OutputModalities, "capabilities": model.Capabilities, + "lifecycle": model.Lifecycle, "regions": model.Regions, "replacement_model": model.ReplacementModel}) } writeJSON(w, map[string]any{"object": "list", "data": data}) } func (a *API) anthropicModels(w http.ResponseWriter, r *http.Request) { - if !a.authorize(w, r) { + principal, ok := a.authorize(w, r) + if !ok { return } - models := a.catalog.Models(domain.ProtocolAnthropic) + models := a.availableModels(a.catalog.ModelsFor(domain.ProtocolAnthropic, principal)) data := make([]map[string]any, 0, len(models)) for _, model := range models { data = append(data, map[string]any{"id": model.ID, "display_name": model.ID, "created_at": "1970-01-01T00:00:00Z", "type": "model"}) @@ -309,13 +352,52 @@ func (a *API) anthropicModels(w http.ResponseWriter, r *http.Request) { writeJSON(w, response) } -func (a *API) authorize(w http.ResponseWriter, r *http.Request) bool { +func (a *API) modelAvailableInRegion(model domain.Model) bool { + if a.deploymentRegion == "" || len(model.Regions) == 0 { + return true + } + for _, region := range model.Regions { + if strings.EqualFold(region, a.deploymentRegion) || region == "*" { + return true + } + } + return false +} + +func (a *API) availableModels(models []domain.Model) []domain.Model { + result := models[:0] + for _, model := range models { + if a.modelAvailableInRegion(model) { + result = append(result, model) + } + } + return result +} +func modelHasCapability(model domain.Model, wanted string) bool { + if len(model.Capabilities) == 0 { + return true + } + for _, value := range model.Capabilities { + if value == wanted || value == "*" { + return true + } + } + return false +} +func modelCreated(model domain.Model) int64 { + if model.ReleasedAt != nil { + return model.ReleasedAt.Unix() + } + return 0 +} + +func (a *API) authorize(w http.ResponseWriter, r *http.Request) (domain.Principal, bool) { principal, err := a.authenticator.Authenticate(r) if err != nil || !hasScope(principal, "inference") { apierror.Write(w, apierror.Error{Status: http.StatusForbidden, Type: "access_denied", Message: "Invalid API key or insufficient permission"}, requestIDFrom(r.Context())) - return false + return domain.Principal{}, false } - return true + return principal, true } func (a *API) publishUsage(event domain.UsageEvent) { @@ -327,11 +409,11 @@ func (a *API) publishUsage(event domain.UsageEvent) { func (a *API) finishUsage(r *http.Request, event domain.UsageEvent) { settled := false if a.billingMeter != nil { - ctx, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), 5*time.Second) - err := a.billingMeter.Settle(ctx, event) + ctx, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), 500*time.Millisecond) + err := a.billingMeter.EnqueueSettlement(ctx, event) cancel() if err != nil { - a.logger.Error("billing_settlement_failed", "request_id", event.RequestID, "tenant_id", event.TenantID, "error", err) + a.logger.Error("billing_settlement_enqueue_failed", "request_id", event.RequestID, "tenant_id", event.TenantID, "error", err) } else { settled = true } diff --git a/internal/httpapi/api_test.go b/internal/httpapi/api_test.go index 014b99b..9407530 100644 --- a/internal/httpapi/api_test.go +++ b/internal/httpapi/api_test.go @@ -37,7 +37,7 @@ func (m *fakeBillingMeter) Authorize(context.Context, billing.Authorization) err return m.authorizeErr } -func (m *fakeBillingMeter) Settle(_ context.Context, event domain.UsageEvent) error { +func (m *fakeBillingMeter) EnqueueSettlement(_ context.Context, event domain.UsageEvent) error { m.settled <- event return nil } diff --git a/internal/httpapi/proxy.go b/internal/httpapi/proxy.go new file mode 100644 index 0000000..3cd2797 --- /dev/null +++ b/internal/httpapi/proxy.go @@ -0,0 +1,44 @@ +package httpapi + +import ( + "net" + "net/http" + "strings" +) + +// TrustProxyHeaders accepts forwarding metadata only from explicitly trusted +// CIDRs. This prevents a direct client from forging HTTPS or audit IP state. +func TrustProxyHeaders(next http.Handler, trustedCIDRs []string, requireHTTPS bool) (http.Handler, error) { + trusted := make([]*net.IPNet, 0, len(trustedCIDRs)) + for _, value := range trustedCIDRs { + _, network, err := net.ParseCIDR(value) + if err != nil { + return nil, err + } + trusted = append(trusted, network) + } + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + host, _, _ := net.SplitHostPort(r.RemoteAddr) + remote := net.ParseIP(host) + trustedPeer := false + for _, network := range trusted { + if remote != nil && network.Contains(remote) { + trustedPeer = true + break + } + } + if !trustedPeer { + for _, header := range []string{"Forwarded", "X-Forwarded-For", "X-Forwarded-Host", "X-Forwarded-Port", "X-Forwarded-Proto", "X-Real-IP"} { + r.Header.Del(header) + } + } else if forwarded := strings.TrimSpace(strings.Split(r.Header.Get("X-Forwarded-For"), ",")[0]); net.ParseIP(forwarded) != nil { + r.RemoteAddr = net.JoinHostPort(forwarded, "0") + } + secure := r.TLS != nil || (trustedPeer && strings.EqualFold(strings.TrimSpace(r.Header.Get("X-Forwarded-Proto")), "https")) + if requireHTTPS && !secure { + http.Error(w, "HTTPS is required", http.StatusUpgradeRequired) + return + } + next.ServeHTTP(w, r) + }), nil +} diff --git a/internal/mailer/mailer.go b/internal/mailer/mailer.go index 4ec11ac..ed6b124 100644 --- a/internal/mailer/mailer.go +++ b/internal/mailer/mailer.go @@ -102,8 +102,9 @@ func (s *Sender) Send(ctx context.Context, message Message) error { } body := strings.ReplaceAll(message.Body, "\r\n", "\n") body = strings.ReplaceAll(body, "\n", "\r\n") - _, writeErr := fmt.Fprintf(w, "From: %s\r\nTo: %s\r\nSubject: %s\r\nMIME-Version: 1.0\r\nContent-Type: text/plain; charset=UTF-8\r\nContent-Transfer-Encoding: 8bit\r\n\r\n%s\r\n", - from, cleanHeader(message.Recipient), cleanHeader(message.Subject), body) + messageID := cleanHeader(message.ID) + "@aigw.local" + _, writeErr := fmt.Fprintf(w, "From: %s\r\nTo: %s\r\nSubject: %s\r\nMessage-ID: <%s>\r\nX-AIGW-Message-ID: %s\r\nMIME-Version: 1.0\r\nContent-Type: text/plain; charset=UTF-8\r\nContent-Transfer-Encoding: 8bit\r\n\r\n%s\r\n", + from, cleanHeader(message.Recipient), cleanHeader(message.Subject), messageID, cleanHeader(message.ID), body) closeErr := w.Close() if writeErr != nil { return fmt.Errorf("write SMTP body: %w", writeErr) diff --git a/internal/operations/operations.go b/internal/operations/operations.go new file mode 100644 index 0000000..e566488 --- /dev/null +++ b/internal/operations/operations.go @@ -0,0 +1,143 @@ +package operations + +import ( + "context" + "encoding/json" + "net/http" + "time" + + "aigw/internal/billing" + "aigw/internal/controlplane" + "aigw/internal/telemetry" +) + +type Handler struct { + Store *controlplane.Store + Manager *controlplane.Manager + Billing *billing.Service + Metrics *telemetry.Metrics + MaxSnapshotAge time.Duration +} + +func (h Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/healthz" { + write(w, http.StatusOK, map[string]any{"status": "alive"}) + return + } + if r.URL.Path == "/metrics" && h.Metrics != nil { + h.Metrics.ServeHTTP(w, r) + return + } + if r.URL.Path != "/readyz" { + http.NotFound(w, r) + return + } + ctx, cancel := context.WithTimeout(r.Context(), 2*time.Second) + defer cancel() + checks := map[string]any{} + ready := true + if h.Store != nil { + if err := h.Store.Ping(ctx); err != nil { + checks["postgres"] = map[string]any{"status": "failed", "error": err.Error()} + ready = false + } else { + checks["postgres"] = map[string]any{"status": "ok"} + } + } + if h.Manager != nil { + healthyAt := h.Manager.LastHealthyAt() + if healthyAt.IsZero() { + healthyAt = h.Manager.LastReloadAt() + } + age := time.Since(healthyAt) + snapshotOK := h.Manager.Loaded() && age <= h.MaxSnapshotAge + checks["snapshot"] = map[string]any{"status": status(snapshotOK), "generation": h.Manager.Generation(), "age_seconds": int64(age.Seconds())} + if !snapshotOK { + ready = false + } + redisStatus := "disabled" + if h.Manager.RedisConfigured() { + redisStatus = "degraded" + if h.Manager.RedisConnected() { + redisStatus = "ok" + } + } + checks["redis"] = map[string]any{"status": redisStatus, "required": false} + } + if h.Billing != nil { + if err := h.Billing.Ping(ctx); err != nil { + checks["billing_postgres"] = map[string]any{"status": "failed", "error": err.Error()} + ready = false + } else { + checks["billing_postgres"] = map[string]any{"status": "ok"} + } + queue, err := h.Billing.SettlementQueueStatus(ctx) + if err != nil { + checks["settlement_queue"] = map[string]any{"status": "failed", "error": err.Error()} + ready = false + } else { + backlog := queue.AwaitingEvent + queue.Pending + queue.Processing + queue.Retrying + queueOK := queue.SpoolRecords == 0 + if queue.OldestPending != nil && time.Since(*queue.OldestPending) > 15*time.Minute { + queueOK = false + } + if !queueOK { + ready = false + } + checks["settlement_queue"] = map[string]any{"status": status(queueOK), "backlog": backlog, "spool_records": queue.SpoolRecords, "oldest_pending": queue.OldestPending} + if h.Metrics != nil { + h.Metrics.SetSettlementQueue(backlog, queue.SpoolRecords) + } + billingHealth, err := h.Billing.OperationalStatus(ctx) + if err != nil { + checks["billing_operations"] = map[string]any{"status": "failed", "error": err.Error()} + ready = false + } else { + billingOK := billingHealth.Ready(time.Now().UTC()) + checks["billing_operations"] = map[string]any{"status": status(billingOK), "details": billingHealth} + if !billingOK { + ready = false + } + if h.Metrics != nil { + h.Metrics.SetStripeOperations(billingHealth) + } + } + } + if h.Store != nil { + mail, err := h.Store.MailQueueStatus(ctx) + if err != nil { + checks["mail_queue"] = map[string]any{"status": "failed", "error": err.Error()} + ready = false + } else { + mailOK := mail.Failed == 0 + if mail.OldestPending != nil && time.Since(*mail.OldestPending) > 10*time.Minute { + mailOK = false + } + checks["mail_queue"] = map[string]any{"status": status(mailOK), "details": mail} + if !mailOK { + ready = false + } + } + } + } + if h.Metrics != nil { + h.Metrics.SetReady(ready) + } + code := http.StatusOK + if !ready { + code = http.StatusServiceUnavailable + } + write(w, code, map[string]any{"status": status(ready), "checks": checks}) +} + +func status(ok bool) string { + if ok { + return "ok" + } + return "failed" +} +func write(w http.ResponseWriter, code int, value any) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(code) + _ = json.NewEncoder(w).Encode(value) +} diff --git a/internal/provider/forwarder.go b/internal/provider/forwarder.go index a9d5734..0d40bb1 100644 --- a/internal/provider/forwarder.go +++ b/internal/provider/forwarder.go @@ -49,7 +49,7 @@ func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, reque if err := ctx.Err(); err != nil { return Result{Attempts: i}, err } - body, err := rewriteModel(originalBody, route.UpstreamModel) + body, err := rewriteRequest(originalBody, route.UpstreamModel, protocol) if err != nil { return Result{Attempts: i}, err } @@ -80,12 +80,29 @@ func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, reque } func rewriteModel(body []byte, upstreamModel string) ([]byte, error) { + return rewriteRequest(body, upstreamModel, "") +} + +func rewriteRequest(body []byte, upstreamModel string, protocol domain.Protocol) ([]byte, error) { var object map[string]json.RawMessage if err := json.Unmarshal(body, &object); err != nil { return nil, fmt.Errorf("decode request body: %w", err) } encoded, _ := json.Marshal(upstreamModel) object["model"] = encoded + if protocol == domain.ProtocolOpenAI { + var stream bool + _ = json.Unmarshal(object["stream"], &stream) + if stream { + var options map[string]json.RawMessage + _ = json.Unmarshal(object["stream_options"], &options) + if options == nil { + options = map[string]json.RawMessage{} + } + options["include_usage"] = json.RawMessage("true") + object["stream_options"], _ = json.Marshal(options) + } + } result, err := json.Marshal(object) if err != nil { return nil, fmt.Errorf("encode upstream request: %w", err) diff --git a/internal/provider/forwarder_test.go b/internal/provider/forwarder_test.go new file mode 100644 index 0000000..2e9b4a1 --- /dev/null +++ b/internal/provider/forwarder_test.go @@ -0,0 +1,39 @@ +package provider + +import ( + "encoding/json" + "testing" + + "aigw/internal/domain" +) + +func TestRewriteRequestForcesOpenAIStreamUsage(t *testing.T) { + result, err := rewriteRequest([]byte(`{"model":"public/model","stream":true,"stream_options":{"other":true}}`), "upstream/model", domain.ProtocolOpenAI) + if err != nil { + t.Fatal(err) + } + var body struct { + Model string `json:"model"` + StreamOptions map[string]any `json:"stream_options"` + } + if err := json.Unmarshal(result, &body); err != nil { + t.Fatal(err) + } + if body.Model != "upstream/model" || body.StreamOptions["include_usage"] != true || body.StreamOptions["other"] != true { + t.Fatalf("unexpected rewritten body: %s", result) + } +} + +func TestRewriteRequestDoesNotAddStreamOptionsToAnthropic(t *testing.T) { + result, err := rewriteRequest([]byte(`{"model":"public/model","stream":true}`), "upstream/model", domain.ProtocolAnthropic) + if err != nil { + t.Fatal(err) + } + var body map[string]json.RawMessage + if err := json.Unmarshal(result, &body); err != nil { + t.Fatal(err) + } + if _, exists := body["stream_options"]; exists { + t.Fatalf("unexpected OpenAI stream options in Anthropic request: %s", result) + } +} diff --git a/internal/security/credentials.go b/internal/security/credentials.go index b55fb6b..900261e 100644 --- a/internal/security/credentials.go +++ b/internal/security/credentials.go @@ -4,54 +4,114 @@ import ( "crypto/aes" "crypto/cipher" "crypto/rand" + "crypto/sha256" "encoding/base64" "errors" "fmt" "io" + "strings" ) type CredentialCipher struct { + current credentialKey + keys []credentialKey +} + +type credentialKey struct { + id []byte aead cipher.AEAD } +var credentialEnvelopeMagic = []byte("AGK1") + func NewCredentialCipher(encodedKey string) (*CredentialCipher, error) { - key, err := base64.StdEncoding.DecodeString(encodedKey) - if err != nil { - return nil, fmt.Errorf("decode credential key: %w", err) - } - if len(key) != 32 { - return nil, errors.New("credential key must be a base64-encoded 32-byte key") - } - block, err := aes.NewCipher(key) - if err != nil { - return nil, fmt.Errorf("create credential cipher: %w", err) + return NewCredentialKeyring(encodedKey, nil) +} + +func NewCredentialKeyring(encodedKey string, previous []string) (*CredentialCipher, error) { + values := append([]string{encodedKey}, previous...) + keys := make([]credentialKey, 0, len(values)) + seen := map[string]struct{}{} + for _, value := range values { + value = strings.TrimSpace(value) + if value == "" { + continue + } + key, err := base64.StdEncoding.DecodeString(value) + if err != nil { + return nil, fmt.Errorf("decode credential key: %w", err) + } + if len(key) != 32 { + return nil, errors.New("credential key must be a base64-encoded 32-byte key") + } + block, err := aes.NewCipher(key) + if err != nil { + return nil, fmt.Errorf("create credential cipher: %w", err) + } + aead, err := cipher.NewGCM(block) + if err != nil { + return nil, fmt.Errorf("create credential AEAD: %w", err) + } + fingerprint := sha256.Sum256(key) + id := base64.RawURLEncoding.EncodeToString(fingerprint[:6]) + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + keys = append(keys, credentialKey{id: []byte(id), aead: aead}) } - aead, err := cipher.NewGCM(block) - if err != nil { - return nil, fmt.Errorf("create credential AEAD: %w", err) + if len(keys) == 0 { + return nil, errors.New("at least one credential key is required") } - return &CredentialCipher{aead: aead}, nil + return &CredentialCipher{current: keys[0], keys: keys}, nil } func (c *CredentialCipher) Encrypt(plaintext string) ([]byte, error) { if plaintext == "" { return nil, errors.New("credential cannot be empty") } - nonce := make([]byte, c.aead.NonceSize()) + nonce := make([]byte, c.current.aead.NonceSize()) if _, err := io.ReadFull(rand.Reader, nonce); err != nil { return nil, fmt.Errorf("generate credential nonce: %w", err) } - return c.aead.Seal(nonce, nonce, []byte(plaintext), nil), nil + prefix := append(append([]byte{}, credentialEnvelopeMagic...), c.current.id...) + sealed := c.current.aead.Seal(nil, nonce, []byte(plaintext), prefix) + return append(append(prefix, nonce...), sealed...), nil } func (c *CredentialCipher) Decrypt(ciphertext []byte) (string, error) { - if len(ciphertext) < c.aead.NonceSize() { - return "", errors.New("credential ciphertext is truncated") + if len(ciphertext) >= len(credentialEnvelopeMagic) && string(ciphertext[:len(credentialEnvelopeMagic)]) == string(credentialEnvelopeMagic) { + idStart := len(credentialEnvelopeMagic) + idEnd := idStart + 8 + if len(ciphertext) < idEnd { + return "", errors.New("credential ciphertext is truncated") + } + prefix := ciphertext[:idEnd] + for _, key := range c.keys { + if string(key.id) != string(ciphertext[idStart:idEnd]) { + continue + } + nonceEnd := idEnd + key.aead.NonceSize() + if len(ciphertext) < nonceEnd { + return "", errors.New("credential ciphertext is truncated") + } + plaintext, err := key.aead.Open(nil, ciphertext[idEnd:nonceEnd], ciphertext[nonceEnd:], prefix) + if err != nil { + return "", errors.New("decrypt credential: authentication failed") + } + return string(plaintext), nil + } + return "", errors.New("credential key is not present in the configured keyring") } - nonce := ciphertext[:c.aead.NonceSize()] - plaintext, err := c.aead.Open(nil, nonce, ciphertext[c.aead.NonceSize():], nil) - if err != nil { - return "", errors.New("decrypt credential: authentication failed") + for _, key := range c.keys { + if len(ciphertext) < key.aead.NonceSize() { + continue + } + nonce := ciphertext[:key.aead.NonceSize()] + plaintext, err := key.aead.Open(nil, nonce, ciphertext[key.aead.NonceSize():], nil) + if err == nil { + return string(plaintext), nil + } } - return string(plaintext), nil + return "", errors.New("decrypt credential: authentication failed") } diff --git a/internal/security/credentials_test.go b/internal/security/credentials_test.go index 07fa59e..570ec50 100644 --- a/internal/security/credentials_test.go +++ b/internal/security/credentials_test.go @@ -2,29 +2,35 @@ package security import ( "encoding/base64" - "strings" "testing" ) -func TestCredentialCipherRoundTrip(t *testing.T) { - key := base64.StdEncoding.EncodeToString([]byte(strings.Repeat("k", 32))) - cipher, err := NewCredentialCipher(key) +func TestCredentialKeyringDecryptsPreviousAndReencryptsWithPrimary(t *testing.T) { + primary := base64.StdEncoding.EncodeToString([]byte("01234567890123456789012345678901")) + previous := base64.StdEncoding.EncodeToString([]byte("abcdefghijklmnopqrstuvwxyzabcdef")) + oldCipher, err := NewCredentialCipher(previous) if err != nil { t.Fatal(err) } - ciphertext, err := cipher.Encrypt("upstream-secret") + legacy, err := oldCipher.Encrypt("provider-secret") if err != nil { t.Fatal(err) } - plaintext, err := cipher.Decrypt(ciphertext) + keyring, err := NewCredentialKeyring(primary, []string{previous}) if err != nil { t.Fatal(err) } - if plaintext != "upstream-secret" { - t.Fatalf("unexpected plaintext: %q", plaintext) + if got, err := keyring.Decrypt(legacy); err != nil || got != "provider-secret" { + t.Fatalf("decrypt previous key: got %q, err %v", got, err) } - ciphertext[len(ciphertext)-1] ^= 1 - if _, err := cipher.Decrypt(ciphertext); err == nil { - t.Fatal("expected authentication failure for modified ciphertext") + rotated, err := keyring.Encrypt("provider-secret") + if err != nil { + t.Fatal(err) + } + if string(rotated[:4]) != "AGK1" { + t.Fatalf("expected versioned credential envelope, got %q", rotated[:4]) + } + if got, err := keyring.Decrypt(rotated); err != nil || got != "provider-secret" { + t.Fatalf("decrypt primary key: got %q, err %v", got, err) } } diff --git a/internal/telemetry/metrics.go b/internal/telemetry/metrics.go index 4942d8d..04cc494 100644 --- a/internal/telemetry/metrics.go +++ b/internal/telemetry/metrics.go @@ -4,14 +4,44 @@ import ( "fmt" "net/http" "sync/atomic" + + "aigw/internal/billing" ) type Metrics struct { - requests atomic.Uint64 - failed atomic.Uint64 - inFlight atomic.Int64 - attempts atomic.Uint64 - droppedUsage atomic.Uint64 + requests atomic.Uint64 + failed atomic.Uint64 + inFlight atomic.Int64 + attempts atomic.Uint64 + droppedUsage atomic.Uint64 + settlementBacklog atomic.Int64 + settlementSpool atomic.Int64 + stripeRefundBacklog atomic.Int64 + stripeUncollected atomic.Int64 + stripeMismatches atomic.Int64 + stripeWebhooks atomic.Int64 + unmeteredSuccesses atomic.Int64 + ready atomic.Int64 +} + +func (m *Metrics) SetSettlementQueue(backlog int64, spool int) { + m.settlementBacklog.Store(backlog) + m.settlementSpool.Store(int64(spool)) +} +func (m *Metrics) SetReady(ready bool) { + if ready { + m.ready.Store(1) + } else { + m.ready.Store(0) + } +} + +func (m *Metrics) SetStripeOperations(status billing.OperationalStatus) { + m.stripeRefundBacklog.Store(status.RefundBacklog) + m.stripeUncollected.Store(status.UncollectedMicros) + m.stripeMismatches.Store(status.ReconciliationMismatches) + m.stripeWebhooks.Store(status.UnprocessedWebhooks) + m.unmeteredSuccesses.Store(status.UnmeteredSuccesses) } func (m *Metrics) RequestStarted() { @@ -41,4 +71,12 @@ func (m *Metrics) ServeHTTP(w http.ResponseWriter, _ *http.Request) { fmt.Fprintf(w, "# TYPE aigw_requests_in_flight gauge\naigw_requests_in_flight %d\n", m.inFlight.Load()) fmt.Fprintf(w, "# TYPE aigw_upstream_attempts_total counter\naigw_upstream_attempts_total %d\n", m.attempts.Load()) fmt.Fprintf(w, "# TYPE aigw_usage_events_dropped_total counter\naigw_usage_events_dropped_total %d\n", m.droppedUsage.Load()) + fmt.Fprintf(w, "# TYPE aigw_billing_settlement_backlog gauge\naigw_billing_settlement_backlog %d\n", m.settlementBacklog.Load()) + fmt.Fprintf(w, "# TYPE aigw_billing_settlement_spool_records gauge\naigw_billing_settlement_spool_records %d\n", m.settlementSpool.Load()) + fmt.Fprintf(w, "# TYPE aigw_stripe_refund_backlog gauge\naigw_stripe_refund_backlog %d\n", m.stripeRefundBacklog.Load()) + fmt.Fprintf(w, "# TYPE aigw_billing_uncollected_micros gauge\naigw_billing_uncollected_micros %d\n", m.stripeUncollected.Load()) + fmt.Fprintf(w, "# TYPE aigw_stripe_reconciliation_mismatches gauge\naigw_stripe_reconciliation_mismatches %d\n", m.stripeMismatches.Load()) + fmt.Fprintf(w, "# TYPE aigw_stripe_webhook_backlog gauge\naigw_stripe_webhook_backlog %d\n", m.stripeWebhooks.Load()) + fmt.Fprintf(w, "# TYPE aigw_billing_unmetered_successes gauge\naigw_billing_unmetered_successes %d\n", m.unmeteredSuccesses.Load()) + fmt.Fprintf(w, "# TYPE aigw_ready gauge\naigw_ready %d\n", m.ready.Load()) } diff --git a/internal/usage/observer.go b/internal/usage/observer.go index cc8b52f..808781f 100644 --- a/internal/usage/observer.go +++ b/internal/usage/observer.go @@ -44,6 +44,17 @@ func (o *Observer) Usage() domain.Usage { return o.usage } +// Reported distinguishes a real zero-token usage object from a response that +// omitted usage entirely. Billing must never infer this from token totals. +func (o *Observer) Reported() bool { + if !o.stream { + o.parseJSON(o.buffer) + } else if len(o.line) > 0 { + o.parseSSELine(o.line) + } + return o.found +} + func (o *Observer) captureTail(p []byte) { if len(p) >= maxCaptureBytes { o.buffer = append(o.buffer[:0], p[len(p)-maxCaptureBytes:]...) diff --git a/internal/usage/observer_test.go b/internal/usage/observer_test.go index 8fcf408..3104acb 100644 --- a/internal/usage/observer_test.go +++ b/internal/usage/observer_test.go @@ -10,6 +10,9 @@ func TestObserverReadsOpenAIJSONUsage(t *testing.T) { observer := NewObserver(domain.ProtocolOpenAI, false) _, _ = observer.Write([]byte(`{"choices":[],"usage":{"prompt_tokens":11,"completion_tokens":7,"total_tokens":18}}`)) got := observer.Usage() + if !observer.Reported() { + t.Fatal("expected usage to be marked as reported") + } if got.InputTokens != 11 || got.OutputTokens != 7 || got.TotalTokens != 18 { t.Fatalf("unexpected usage: %+v", got) } @@ -20,7 +23,23 @@ func TestObserverCombinesAnthropicSSEUsage(t *testing.T) { _, _ = observer.Write([]byte("event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":12,\"output_tokens\":1}}}\n\n")) _, _ = observer.Write([]byte("event: message_delta\ndata: {\"type\":\"message_delta\",\"usage\":{\"output_tokens\":8}}\n\n")) got := observer.Usage() + if !observer.Reported() { + t.Fatal("expected streaming usage to be marked as reported") + } if got.InputTokens != 12 || got.OutputTokens != 8 || got.TotalTokens != 20 { t.Fatalf("unexpected usage: %+v", got) } } + +func TestObserverDistinguishesMissingUsageFromReportedZero(t *testing.T) { + missing := NewObserver(domain.ProtocolOpenAI, false) + _, _ = missing.Write([]byte(`{"choices":[]}`)) + if missing.Reported() { + t.Fatal("response without usage must not be reported") + } + reported := NewObserver(domain.ProtocolOpenAI, false) + _, _ = reported.Write([]byte(`{"choices":[],"usage":{"prompt_tokens":0,"completion_tokens":0,"total_tokens":0}}`)) + if !reported.Reported() { + t.Fatal("explicit zero usage must be distinguished from a missing usage object") + } +} |
