refactor pricing, provider fees, realtime model map updates

This commit is contained in:
Shroominic
2025-10-20 12:45:02 +08:00
parent 61a0559f8e
commit 0da08fb945
11 changed files with 854 additions and 424 deletions
@@ -0,0 +1,64 @@
"""change models to composite primary key (id, upstream_provider_id)
Revision ID: a1a1a1a1a1a1
Revises: f7a8b9c0d1e2
Create Date: 2025-10-20 00:00:00.000000
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
revision = "a1a1a1a1a1a1"
down_revision = "f7a8b9c0d1e2"
branch_labels = None
depends_on = None
def upgrade() -> None:
conn = op.get_bind()
inspector = sa.inspect(conn)
if "models" in inspector.get_table_names():
op.drop_table("models")
op.create_table(
"models",
sa.Column("id", sa.String(), nullable=False),
sa.Column("upstream_provider_id", sa.Integer(), nullable=False),
sa.Column("name", sa.String(), nullable=False),
sa.Column("created", sa.Integer(), nullable=False),
sa.Column("description", sa.Text(), nullable=False),
sa.Column("context_length", sa.Integer(), nullable=False),
sa.Column("architecture", sa.Text(), nullable=False),
sa.Column("pricing", sa.Text(), nullable=False),
sa.Column("sats_pricing", sa.Text(), nullable=True),
sa.Column("per_request_limits", sa.Text(), nullable=True),
sa.Column("top_provider", sa.Text(), nullable=True),
sa.Column("enabled", sa.Boolean(), nullable=False, server_default="1"),
sa.PrimaryKeyConstraint("id", "upstream_provider_id"),
sa.ForeignKeyConstraint(
["upstream_provider_id"], ["upstream_providers.id"], ondelete="CASCADE"
),
)
def downgrade() -> None:
op.drop_table("models")
op.create_table(
"models",
sa.Column("id", sa.String(), primary_key=True, nullable=False),
sa.Column("name", sa.String(), nullable=False),
sa.Column("created", sa.Integer(), nullable=False),
sa.Column("description", sa.Text(), nullable=False),
sa.Column("context_length", sa.Integer(), nullable=False),
sa.Column("architecture", sa.Text(), nullable=False),
sa.Column("pricing", sa.Text(), nullable=False),
sa.Column("sats_pricing", sa.Text(), nullable=True),
sa.Column("per_request_limits", sa.Text(), nullable=True),
sa.Column("top_provider", sa.Text(), nullable=True),
sa.Column("enabled", sa.Boolean(), nullable=False, server_default="1"),
sa.Column("upstream_provider_id", sa.Integer(), nullable=True),
sa.ForeignKeyConstraint(["upstream_provider_id"], ["upstream_providers.id"]),
)
@@ -0,0 +1,27 @@
"""add provider_fee to upstream_providers
Revision ID: f7a8b9c0d1e2
Revises: e1f2a3b4c5d6
Create Date: 2025-10-13 00:00:00.000000
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
revision = "f7a8b9c0d1e2"
down_revision = "e1f2a3b4c5d6"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"upstream_providers",
sa.Column("provider_fee", sa.Float(), nullable=False, server_default="1.01"),
)
def downgrade() -> None:
op.drop_column("upstream_providers", "provider_fee")
+353 -145
View File
@@ -1389,30 +1389,30 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
<script>
let providersList = [];
let selectedProviderId = null;
let providerModels = { db_models: [], remote_models: [] };
let providerModels = { db_models: [], remote_models: [], provider: {} };
let openrouterPresets = [];
async function fetchProviders() {
const tableBody = document.getElementById('providers-tbody');
tableBody.innerHTML = '<tr><td colspan="5" style="color:#718096;">Loading…</td></tr>';
tableBody.innerHTML = '<tr><td colspan="4" style="color:#718096;">Loading…</td></tr>';
try {
const resp = await fetch('/admin/api/upstream-providers', { credentials: 'same-origin' });
if (!resp.ok) throw new Error('HTTP ' + resp.status);
providersList = await resp.json();
renderProvidersTable();
} catch (e) {
tableBody.innerHTML = '<tr><td colspan="5" style="color:#e53e3e;">Failed to load providers: ' + e.message + '</td></tr>';
tableBody.innerHTML = '<tr><td colspan="4" style="color:#e53e3e;">Failed to load providers: ' + e.message + '</td></tr>';
}
}
function renderProvidersTable() {
const tableBody = document.getElementById('providers-tbody');
if (!Array.isArray(providersList) || !providersList.length) {
tableBody.innerHTML = '<tr><td colspan="5" style="color:#718096;">No providers found</td></tr>';
tableBody.innerHTML = '<tr><td colspan="4" style="color:#718096;">No providers found</td></tr>';
return;
}
const rows = providersList.map(p => `
<tr>
<td>${p.id}</td>
<td>${p.provider_type}</td>
<td style="word-break: break-all;">${p.base_url}</td>
<td><span style="padding:2px 8px; border-radius:4px; background:${p.enabled ? '#22c55e' : '#ef4444'}; color:white; font-size:12px;">${p.enabled ? 'Enabled' : 'Disabled'}</span></td>
@@ -1455,6 +1455,9 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
function renderProviderModels() {
const dbModelsBody = document.getElementById('db-models-tbody');
const remoteModelsBody = document.getElementById('remote-models-tbody');
const remoteModelsSection = document.getElementById('remote-models-section');
const customProviderActions = document.getElementById('custom-provider-actions');
const isCustomProvider = providerModels.provider && providerModels.provider.provider_type === 'custom';
if (providerModels.db_models && providerModels.db_models.length > 0) {
dbModelsBody.innerHTML = '';
@@ -1483,32 +1486,40 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
dbModelsBody.appendChild(row);
});
} else {
dbModelsBody.innerHTML = '<tr><td colspan="3" style="color:#718096;">No model overrides</td></tr>';
dbModelsBody.innerHTML = '<tr><td colspan="3" style="color:#718096;">No models defined</td></tr>';
}
if (providerModels.remote_models && providerModels.remote_models.length > 0) {
remoteModelsBody.innerHTML = '';
providerModels.remote_models.forEach(m => {
const isInDb = providerModels.db_models.some(db => db.id === m.id);
const row = document.createElement('tr');
row.innerHTML = `
<td style="font-family:monospace; word-break: break-all;">${m.id}</td>
<td>${m.name}</td>
<td>
<button class="override-btn" ${isInDb ? 'disabled' : ''}>+ Override</button>
<button class="disable-btn" style="background:#ef4444;" ${isInDb ? 'disabled' : ''}>🚫 Disable</button>
</td>
`;
const overrideBtn = row.querySelector('.override-btn');
const disableBtn = row.querySelector('.disable-btn');
if (!isInDb) {
overrideBtn.onclick = () => createModelOverride(m, true);
disableBtn.onclick = () => createModelOverride(m, false);
}
remoteModelsBody.appendChild(row);
});
if (isCustomProvider) {
remoteModelsSection.style.display = 'none';
customProviderActions.style.display = 'block';
} else {
remoteModelsBody.innerHTML = '<tr><td colspan="3" style="color:#718096;">No remote models available</td></tr>';
remoteModelsSection.style.display = 'block';
customProviderActions.style.display = 'none';
if (providerModels.remote_models && providerModels.remote_models.length > 0) {
remoteModelsBody.innerHTML = '';
providerModels.remote_models.forEach(m => {
const isInDb = providerModels.db_models.some(db => db.id === m.id);
const row = document.createElement('tr');
row.innerHTML = `
<td style="font-family:monospace; word-break: break-all;">${m.id}</td>
<td>${m.name}</td>
<td>
<button class="override-btn" ${isInDb ? 'disabled' : ''}>+ Override</button>
<button class="disable-btn" style="background:#ef4444;" ${isInDb ? 'disabled' : ''}>🚫 Disable</button>
</td>
`;
const overrideBtn = row.querySelector('.override-btn');
const disableBtn = row.querySelector('.disable-btn');
if (!isInDb) {
overrideBtn.onclick = () => createModelOverride(m, true);
disableBtn.onclick = () => createModelOverride(m, false);
}
remoteModelsBody.appendChild(row);
});
} else {
remoteModelsBody.innerHTML = '<tr><td colspan="3" style="color:#718096;">No remote models available</td></tr>';
}
}
}
@@ -1522,15 +1533,20 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
'openrouter': { baseUrl: 'https://openrouter.ai/api/v1', showApiVersion: false },
'anthropic': { baseUrl: 'https://api.anthropic.com/v1', showApiVersion: false },
'azure': { baseUrl: '', showApiVersion: true },
'generic': { baseUrl: '', showApiVersion: false }
'custom': { baseUrl: '', showApiVersion: false }
};
function getProviderFeePlaceholder(providerType) {
return providerType === 'openrouter' ? 'Default: 1.06 (6%)' : 'Default: 1.01 (1%)';
}
function updateProviderFields() {
const providerType = document.getElementById('provider-type').value;
const config = PROVIDER_CONFIGS[providerType] || { baseUrl: '', showApiVersion: false };
const baseUrlField = document.getElementById('provider-base-url');
const apiVersionRow = document.getElementById('api-version-row');
const providerId = document.getElementById('provider-id').value;
const feeField = document.getElementById('provider-fee');
if (config.baseUrl && !providerId) {
baseUrlField.value = config.baseUrl;
@@ -1542,6 +1558,10 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
}
apiVersionRow.style.display = config.showApiVersion ? 'block' : 'none';
if (!feeField.value) {
feeField.placeholder = getProviderFeePlaceholder(providerType);
}
}
async function openProviderEditor(providerId) {
@@ -1567,6 +1587,7 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
document.getElementById('provider-api-key').placeholder = '[Keep existing]';
document.getElementById('provider-api-version').value = p.api_version || '';
document.getElementById('provider-enabled').checked = p.enabled;
document.getElementById('provider-fee').value = p.provider_fee || '';
updateProviderFields();
} catch (e) {
errorBox.style.display = 'block';
@@ -1581,6 +1602,8 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
document.getElementById('provider-api-key').placeholder = 'API Key';
document.getElementById('provider-api-version').value = '';
document.getElementById('provider-enabled').checked = true;
document.getElementById('provider-fee').value = '';
document.getElementById('provider-fee').placeholder = getProviderFeePlaceholder('openrouter');
updateProviderFields();
}
@@ -1596,11 +1619,16 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
const errorBox = document.getElementById('provider-error');
errorBox.style.display = 'none';
const providerType = document.getElementById('provider-type').value;
const feeValue = document.getElementById('provider-fee').value;
const defaultFee = providerType === 'openrouter' ? 1.06 : 1.01;
const payload = {
provider_type: document.getElementById('provider-type').value,
provider_type: providerType,
base_url: document.getElementById('provider-base-url').value,
api_version: document.getElementById('provider-api-version').value || null,
enabled: document.getElementById('provider-enabled').checked,
provider_fee: feeValue ? parseFloat(feeValue) : defaultFee,
};
const apiKey = document.getElementById('provider-api-key').value;
@@ -1675,13 +1703,128 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
}
}
async function createModelOverride(modelData, enabled) {
async function openCustomModelCreator() {
const modal = document.getElementById('model-override-modal');
const errorBox = document.getElementById('model-override-error');
errorBox.style.display = 'none';
document.getElementById('override-model-id').value = '';
document.getElementById('override-model-id').disabled = false;
document.getElementById('override-model-name').value = '';
document.getElementById('override-description').value = '';
document.getElementById('override-context').value = 8192;
const architecture = {
modality: 'text',
input_modalities: ['text'],
output_modalities: ['text'],
tokenizer: '',
instruct_type: null
};
document.getElementById('override-architecture').value = JSON.stringify(architecture, null, 2);
const pricing = {
prompt: 0.0,
completion: 0.0,
request: 0.0,
image: 0.0,
web_search: 0.0,
internal_reasoning: 0.0
};
document.getElementById('override-pricing').value = JSON.stringify(pricing, null, 2);
document.getElementById('override-enabled').value = 'true';
document.getElementById('override-mode').value = 'create';
document.getElementById('override-created').value = Math.floor(Date.now() / 1000);
document.getElementById('override-upstream-provider-id').value = selectedProviderId;
document.getElementById('modal-title').textContent = 'Create Custom Model';
document.getElementById('override-save-btn').textContent = 'Create Model';
document.getElementById('override-model-name').disabled = false;
modal.style.display = 'block';
}
async function openPresetSelector() {
const modal = document.getElementById('preset-selector-modal');
const errorBox = document.getElementById('preset-error');
const searchInput = document.getElementById('preset-search');
errorBox.style.display = 'none';
searchInput.value = '';
if (!openrouterPresets.length) {
const loadingDiv = document.getElementById('preset-loading');
const presetsDiv = document.getElementById('presets-list');
loadingDiv.style.display = 'block';
presetsDiv.style.display = 'none';
try {
const resp = await fetch('/admin/api/openrouter-presets', { credentials: 'same-origin' });
if (!resp.ok) throw new Error('HTTP ' + resp.status);
openrouterPresets = await resp.json();
renderPresets('');
loadingDiv.style.display = 'none';
presetsDiv.style.display = 'block';
} catch (e) {
errorBox.style.display = 'block';
errorBox.textContent = 'Failed to load presets: ' + e.message;
loadingDiv.style.display = 'none';
}
} else {
renderPresets('');
}
modal.style.display = 'block';
}
function renderPresets(query) {
const presetsBody = document.getElementById('presets-tbody');
const q = (query || '').trim().toLowerCase();
const filtered = q ? openrouterPresets.filter(m => {
const id = (m.id || '').toLowerCase();
const name = (m.name || '').toLowerCase();
return id.includes(q) || name.includes(q);
}) : openrouterPresets;
if (!filtered.length) {
presetsBody.innerHTML = '<tr><td colspan="3" style="color:#718096;">No models match your search</td></tr>';
return;
}
presetsBody.innerHTML = '';
filtered.slice(0, 100).forEach(m => {
const row = document.createElement('tr');
row.innerHTML = `
<td style="font-family:monospace; word-break: break-all; font-size: 0.85rem;">${m.id}</td>
<td>${m.name}</td>
<td><button class="use-preset-btn">Use Preset</button></td>
`;
const btn = row.querySelector('.use-preset-btn');
btn.onclick = () => usePreset(m);
presetsBody.appendChild(row);
});
}
function searchPresets(query) {
renderPresets(query);
}
function closePresetSelector() {
document.getElementById('preset-selector-modal').style.display = 'none';
}
async function usePreset(modelData) {
closePresetSelector();
await createModelOverride(modelData, true, true);
}
async function createModelOverride(modelData, enabled, isCustomModel = false) {
if (enabled) {
const modal = document.getElementById('model-override-modal');
const errorBox = document.getElementById('model-override-error');
errorBox.style.display = 'none';
document.getElementById('override-model-id').value = modelData.id || '';
document.getElementById('override-model-id').disabled = !isCustomModel;
document.getElementById('override-model-name').value = modelData.name || '';
document.getElementById('override-description').value = modelData.description || '';
document.getElementById('override-context').value = modelData.context_length || 8192;
@@ -1699,9 +1842,10 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
prompt: 0.0,
completion: 0.0,
request: 0.0,
image: 0.0
image: 0.0,
web_search: 0.0,
internal_reasoning: 0.0
};
// Remove computed fields
delete pricing.max_prompt_cost;
delete pricing.max_completion_cost;
delete pricing.max_cost;
@@ -1711,8 +1855,9 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
document.getElementById('override-mode').value = 'create';
document.getElementById('override-created').value = Math.floor(Date.now() / 1000);
document.getElementById('override-upstream-provider-id').value = selectedProviderId;
document.getElementById('modal-title').textContent = 'Create Model Override';
document.getElementById('override-save-btn').textContent = 'Create Override';
const isCustomProvider = providerModels.provider && providerModels.provider.provider_type === 'custom';
document.getElementById('modal-title').textContent = isCustomProvider ? 'Create Model from Preset' : 'Create Model Override';
document.getElementById('override-save-btn').textContent = isCustomProvider ? 'Create Model' : 'Create Override';
document.getElementById('override-model-name').disabled = false;
modal.style.display = 'block';
@@ -1735,7 +1880,7 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
enabled: false
};
const resp = await fetch('/admin/api/models', {
const resp = await fetch(`/admin/api/upstream-providers/${selectedProviderId}/models`, {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
credentials: 'same-origin',
@@ -1756,7 +1901,7 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
async function toggleModelEnabled(modelId, newEnabledState) {
try {
const resp = await fetch(`/admin/api/models/${encodeURIComponent(modelId)}`, {
const resp = await fetch(`/admin/api/upstream-providers/${selectedProviderId}/models/${encodeURIComponent(modelId)}`, {
credentials: 'same-origin'
});
if (!resp.ok) throw new Error('Failed to fetch model');
@@ -1764,7 +1909,7 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
model.enabled = newEnabledState;
const updateResp = await fetch(`/admin/api/models/${encodeURIComponent(modelId)}`, {
const updateResp = await fetch(`/admin/api/upstream-providers/${selectedProviderId}/models/${encodeURIComponent(modelId)}`, {
method: 'PATCH',
headers: { 'Content-Type': 'application/json' },
credentials: 'same-origin',
@@ -1809,7 +1954,7 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
};
const resp = await fetch(
isEdit ? `/admin/api/models/${encodeURIComponent(modelId)}` : '/admin/api/models',
isEdit ? `/admin/api/upstream-providers/${selectedProviderId}/models/${encodeURIComponent(modelId)}` : `/admin/api/upstream-providers/${selectedProviderId}/models`,
{
method: isEdit ? 'PATCH' : 'POST',
headers: { 'Content-Type': 'application/json' },
@@ -1841,7 +1986,7 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
errorBox.style.display = 'none';
try {
const resp = await fetch(`/admin/api/models/${encodeURIComponent(modelId)}`, {
const resp = await fetch(`/admin/api/upstream-providers/${selectedProviderId}/models/${encodeURIComponent(modelId)}`, {
credentials: 'same-origin'
});
if (!resp.ok) throw new Error('Failed to fetch model');
@@ -1891,7 +2036,7 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
if (!confirm('Delete this model override?')) return;
try {
const resp = await fetch(`/admin/api/models/${encodeURIComponent(modelId)}`, {
const resp = await fetch(`/admin/api/upstream-providers/${selectedProviderId}/models/${encodeURIComponent(modelId)}`, {
method: 'DELETE',
credentials: 'same-origin'
});
@@ -1907,8 +2052,10 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
window.onclick = function(event) {
const editModal = document.getElementById('provider-edit-modal');
const overrideModal = document.getElementById('model-override-modal');
const presetModal = document.getElementById('preset-selector-modal');
if (event.target == editModal) closeProviderEditor();
if (event.target == overrideModal) closeModelOverrideModal();
if (event.target == presetModal) closePresetSelector();
}
</script>
"""
@@ -1936,7 +2083,6 @@ def upstream_providers_page() -> str:
<table>
<thead>
<tr>
<th>ID</th>
<th>Type</th>
<th>Base URL</th>
<th>Status</th>
@@ -1944,7 +2090,7 @@ def upstream_providers_page() -> str:
</tr>
</thead>
<tbody id="providers-tbody">
<tr><td colspan="5" style="color:#718096;">Loading…</td></tr>
<tr><td colspan="4" style="color:#718096;">Loading…</td></tr>
</tbody>
</table>
</div>
@@ -1958,7 +2104,15 @@ def upstream_providers_page() -> str:
<div id="models-loading" style="color:#718096;">Loading models…</div>
<div id="models-content" style="display:none;">
<h3>Database Overrides</h3>
<div id="custom-provider-actions" style="display:none; margin-bottom: 1rem;">
<button onclick="openCustomModelCreator()"> Add Model</button>
<button onclick="openPresetSelector()" style="background:#48bb78;">📋 Load from Preset</button>
<p style="font-size: 0.9rem; color: #718096; margin-top: 0.5rem;">
Custom providers don't fetch models from an API. Add models manually or use OpenRouter presets.
</p>
</div>
<h3>Models</h3>
<table style="margin-bottom: 2rem;">
<thead>
<tr>
@@ -1968,23 +2122,25 @@ def upstream_providers_page() -> str:
</tr>
</thead>
<tbody id="db-models-tbody">
<tr><td colspan="3" style="color:#718096;">No overrides</td></tr>
<tr><td colspan="3" style="color:#718096;">No models defined</td></tr>
</tbody>
</table>
<h3>Remote Models</h3>
<table>
<thead>
<tr>
<th>Model ID</th>
<th>Name</th>
<th>Actions</th>
</tr>
</thead>
<tbody id="remote-models-tbody">
<tr><td colspan="3" style="color:#718096;">No remote models</td></tr>
</tbody>
</table>
<div id="remote-models-section">
<h3>Remote Models</h3>
<table>
<thead>
<tr>
<th>Model ID</th>
<th>Name</th>
<th>Actions</th>
</tr>
</thead>
<tbody id="remote-models-tbody">
<tr><td colspan="3" style="color:#718096;">No remote models</td></tr>
</tbody>
</table>
</div>
</div>
</div>
@@ -2003,7 +2159,7 @@ def upstream_providers_page() -> str:
<option value="openai">OpenAI</option>
<option value="anthropic">Anthropic</option>
<option value="azure">Azure OpenAI</option>
<option value="generic">Generic</option>
<option value="custom">Custom</option>
</select>
<label>Base URL</label>
@@ -2017,6 +2173,10 @@ def upstream_providers_page() -> str:
<input type="text" id="provider-api-version" placeholder="2024-02-15-preview">
</div>
<label>Provider Fee (Multiplier)</label>
<input type="number" id="provider-fee" step="0.001" min="1.0" placeholder="Default: 1.01 (1%)">
<small style="color:#718096; font-size:0.85rem;">Leave empty to use default (OpenRouter: 1.06, Others: 1.01)</small>
<label style="display:flex; align-items:center; gap:8px; margin:10px 0;">
<input type="checkbox" id="provider-enabled" style="width:auto;">
<span>Enabled</span>
@@ -2065,6 +2225,48 @@ def upstream_providers_page() -> str:
</div>
</div>
</div>
<div id="preset-selector-modal" class="modal">
<div class="modal-content" style="max-width: 900px;">
<span class="close" onclick="closePresetSelector()">&times;</span>
<h3>Load Model from OpenRouter Preset</h3>
<div id="preset-error" style="display:none; margin: 10px 0; color:#e53e3e;"></div>
<div id="preset-loading" style="color:#718096; padding: 20px; text-align: center;">
Loading OpenRouter models...
</div>
<div id="presets-list" style="display:none;">
<input type="text" id="preset-search" placeholder="Search models by ID or name..."
oninput="searchPresets(this.value)"
style="width:100%; padding:10px; margin-bottom:12px; border:2px solid #e2e8f0; border-radius:6px;">
<div style="max-height: 60vh; overflow-y: auto;">
<table style="margin: 0;">
<thead style="position: sticky; top: 0; z-index: 10;">
<tr>
<th style="width: 35%;">Model ID</th>
<th style="width: 45%;">Name</th>
<th style="width: 20%;">Actions</th>
</tr>
</thead>
<tbody id="presets-tbody">
<tr><td colspan="3" style="color:#718096;">Loading...</td></tr>
</tbody>
</table>
</div>
<p style="font-size: 0.85rem; color: #718096; margin-top: 10px;">
Showing up to 100 models. Use search to find specific models.
</p>
</div>
<div style="margin-top: 12px;">
<button onclick="closePresetSelector()" style="background-color: #718096;">Cancel</button>
</div>
</div>
</div>
</body>
</html>
"""
@@ -2078,12 +2280,6 @@ async def admin_upstream_providers(request: Request) -> str:
return admin_auth()
@admin_router.get("/api/models", dependencies=[Depends(require_admin_api)])
async def get_models_admin_api(request: Request) -> list[dict[str, object]]:
items = await list_models()
return [m.dict() for m in items] # type: ignore
class ModelCreate(BaseModel):
id: str
name: str
@@ -2098,14 +2294,25 @@ class ModelCreate(BaseModel):
enabled: bool = True
@admin_router.post("/api/models", dependencies=[Depends(require_admin_api)])
async def create_model_admin_api(payload: ModelCreate) -> dict[str, object]:
@admin_router.post(
"/api/upstream-providers/{provider_id}/models",
dependencies=[Depends(require_admin_api)],
)
async def create_provider_model(
provider_id: int, payload: ModelCreate
) -> dict[str, object]:
async with create_session() as session:
exists = await session.get(ModelRow, payload.id)
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
exists = await session.get(ModelRow, (payload.id, provider_id))
if exists:
raise HTTPException(
status_code=409, detail="Model with this ID already exists"
status_code=409,
detail="Model with this ID already exists for this provider",
)
row = ModelRow(
id=payload.id,
name=payload.name,
@@ -2123,7 +2330,7 @@ async def create_model_admin_api(payload: ModelCreate) -> dict[str, object]:
top_provider=(
json.dumps(payload.top_provider) if payload.top_provider else None
),
upstream_provider_id=payload.upstream_provider_id,
upstream_provider_id=provider_id,
enabled=payload.enabled,
)
session.add(row)
@@ -2131,70 +2338,29 @@ async def create_model_admin_api(payload: ModelCreate) -> dict[str, object]:
await session.refresh(row)
await refresh_model_maps()
return _row_to_model(row).dict() # type: ignore
@admin_router.post("/api/models/batch", dependencies=[Depends(require_admin_api)])
async def batch_create_models(payload: dict[str, object]) -> dict[str, int]:
models = payload.get("models")
if not isinstance(models, list) or not models:
raise HTTPException(
status_code=400, detail="Payload must include non-empty 'models' array"
)
created = 0
skipped = 0
async with create_session() as session:
for m in models:
try:
model = Model(**m) # type: ignore[arg-type]
except Exception:
skipped += 1
continue
exists = await session.get(ModelRow, model.id)
if exists:
skipped += 1
continue
pricing_dict = model.pricing.dict()
for k in ("max_prompt_cost", "max_completion_cost", "max_cost"):
pricing_dict.pop(k, None)
row = ModelRow(
id=model.id,
name=model.name,
description=model.description,
created=int(model.created),
context_length=int(model.context_length),
architecture=json.dumps(model.architecture.dict()),
pricing=json.dumps(pricing_dict),
sats_pricing=None,
per_request_limits=(
json.dumps(model.per_request_limits)
if model.per_request_limits is not None
else None
),
top_provider=(
json.dumps(model.top_provider.dict())
if model.top_provider
else None
),
)
session.add(row)
created += 1
if created:
await session.commit()
if created:
await refresh_model_maps()
return {"created": created, "skipped": skipped}
return _row_to_model(
row, apply_provider_fee=True, provider_fee=provider.provider_fee
).dict() # type: ignore
@admin_router.get(
"/api/models/{model_id:path}", dependencies=[Depends(require_admin_api)]
"/api/upstream-providers/{provider_id}/models/{model_id:path}",
dependencies=[Depends(require_admin_api)],
)
async def get_model_admin_api(model_id: str) -> dict[str, object]:
async def get_provider_model(provider_id: int, model_id: str) -> dict[str, object]:
async with create_session() as session:
row = await session.get(ModelRow, model_id)
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
row = await session.get(ModelRow, (model_id, provider_id))
if not row:
raise HTTPException(status_code=404, detail="Model not found")
return _row_to_model(row).dict() # type: ignore
raise HTTPException(
status_code=404, detail="Model not found for this provider"
)
return _row_to_model(
row, apply_provider_fee=True, provider_fee=provider.provider_fee
).dict() # type: ignore
class ModelUpdate(BaseModel):
@@ -2212,18 +2378,25 @@ class ModelUpdate(BaseModel):
@admin_router.patch(
"/api/models/{model_id:path}", dependencies=[Depends(require_admin_api)]
"/api/upstream-providers/{provider_id}/models/{model_id:path}",
dependencies=[Depends(require_admin_api)],
)
async def update_model_admin_api(
model_id: str, payload: ModelUpdate
async def update_provider_model(
provider_id: int, model_id: str, payload: ModelUpdate
) -> dict[str, object]:
if payload.id != model_id:
raise HTTPException(status_code=400, detail="Path id does not match payload id")
async with create_session() as session:
row = await session.get(ModelRow, model_id)
provider = await session.get(UpstreamProviderRow, provider_id)
if not provider:
raise HTTPException(status_code=404, detail="Provider not found")
row = await session.get(ModelRow, (model_id, provider_id))
if not row:
raise HTTPException(status_code=404, detail="Model not found")
raise HTTPException(
status_code=404, detail="Model not found for this provider"
)
row.name = payload.name
row.description = payload.description
@@ -2240,7 +2413,6 @@ async def update_model_admin_api(
row.top_provider = (
json.dumps(payload.top_provider) if payload.top_provider else None
)
row.upstream_provider_id = payload.upstream_provider_id
row.enabled = payload.enabled
session.add(row)
@@ -2248,33 +2420,43 @@ async def update_model_admin_api(
await session.refresh(row)
await refresh_model_maps()
return _row_to_model(row).dict() # type: ignore
return _row_to_model(
row, apply_provider_fee=True, provider_fee=provider.provider_fee
).dict() # type: ignore
@admin_router.delete(
"/api/models/{model_id:path}", dependencies=[Depends(require_admin_api)]
"/api/upstream-providers/{provider_id}/models/{model_id:path}",
dependencies=[Depends(require_admin_api)],
)
async def delete_model_admin_api(model_id: str) -> dict[str, object]:
async def delete_provider_model(provider_id: int, model_id: str) -> dict[str, object]:
async with create_session() as session:
row = await session.get(ModelRow, model_id)
row = await session.get(ModelRow, (model_id, provider_id))
if not row:
raise HTTPException(status_code=404, detail="Model not found")
raise HTTPException(
status_code=404, detail="Model not found for this provider"
)
await session.delete(row)
await session.commit()
await refresh_model_maps()
return {"ok": True, "deleted_id": model_id}
@admin_router.delete("/api/models", dependencies=[Depends(require_admin_api)])
async def delete_all_models_admin_api() -> dict[str, object]:
@admin_router.delete(
"/api/upstream-providers/{provider_id}/models",
dependencies=[Depends(require_admin_api)],
)
async def delete_all_provider_models(provider_id: int) -> dict[str, object]:
async with create_session() as session:
result = await session.exec(select(ModelRow)) # type: ignore
result = await session.exec(
select(ModelRow).where(ModelRow.upstream_provider_id == provider_id)
) # type: ignore
rows = result.all()
for row in rows:
await session.delete(row) # type: ignore
await session.commit()
await refresh_model_maps()
return {"ok": True, "deleted": "all"}
return {"ok": True, "deleted": len(rows)}
class UpstreamProviderCreate(BaseModel):
@@ -2283,6 +2465,7 @@ class UpstreamProviderCreate(BaseModel):
api_key: str
api_version: str | None = None
enabled: bool = True
provider_fee: float = 1.01
class UpstreamProviderUpdate(BaseModel):
@@ -2291,6 +2474,7 @@ class UpstreamProviderUpdate(BaseModel):
api_key: str | None = None
api_version: str | None = None
enabled: bool | None = None
provider_fee: float | None = None
@admin_router.get("/api/upstream-providers", dependencies=[Depends(require_admin_api)])
@@ -2306,6 +2490,7 @@ async def get_upstream_providers() -> list[dict[str, object]]:
"api_key": "[REDACTED]" if p.api_key else "",
"api_version": p.api_version,
"enabled": p.enabled,
"provider_fee": p.provider_fee,
}
for p in providers
]
@@ -2332,12 +2517,14 @@ async def create_upstream_provider(
api_key=payload.api_key,
api_version=payload.api_version,
enabled=payload.enabled,
provider_fee=payload.provider_fee,
)
session.add(provider)
await session.commit()
await session.refresh(provider)
await reinitialize_upstreams()
await refresh_model_maps()
return {
"id": provider.id,
"provider_type": provider.provider_type,
@@ -2345,6 +2532,7 @@ async def create_upstream_provider(
"api_key": "[REDACTED]",
"api_version": provider.api_version,
"enabled": provider.enabled,
"provider_fee": provider.provider_fee,
}
@@ -2363,6 +2551,7 @@ async def get_upstream_provider(provider_id: int) -> dict[str, object]:
"api_key": "[REDACTED]" if provider.api_key else "",
"api_version": provider.api_version,
"enabled": provider.enabled,
"provider_fee": provider.provider_fee,
}
@@ -2387,12 +2576,15 @@ async def update_upstream_provider(
provider.api_version = payload.api_version
if payload.enabled is not None:
provider.enabled = payload.enabled
if payload.provider_fee is not None:
provider.provider_fee = payload.provider_fee
session.add(provider)
await session.commit()
await session.refresh(provider)
await reinitialize_upstreams()
await refresh_model_maps()
return {
"id": provider.id,
"provider_type": provider.provider_type,
@@ -2400,6 +2592,7 @@ async def update_upstream_provider(
"api_key": "[REDACTED]",
"api_version": provider.api_version,
"enabled": provider.enabled,
"provider_fee": provider.provider_fee,
}
@@ -2414,6 +2607,7 @@ async def delete_upstream_provider(provider_id: int) -> dict[str, object]:
await session.delete(provider)
await session.commit()
await reinitialize_upstreams()
await refresh_model_maps()
return {"ok": True, "deleted_id": provider_id}
@@ -2437,8 +2631,11 @@ async def get_provider_models(provider_id: int) -> dict[str, object]:
upstream_instance = _instantiate_provider(provider)
if upstream_instance:
try:
models = await upstream_instance.fetch_models()
remote_models = [m.dict() for m in models]
raw_models = await upstream_instance.fetch_models()
remote_models = [
upstream_instance._apply_provider_fee_to_model(m).dict()
for m in raw_models
]
except Exception as e:
logger.error(
f"Failed to fetch models from {provider.provider_type}: {e}"
@@ -2455,6 +2652,17 @@ async def get_provider_models(provider_id: int) -> dict[str, object]:
}
@admin_router.get(
"/api/openrouter-presets",
dependencies=[Depends(require_admin_api)],
)
async def get_openrouter_presets() -> list[dict[str, object]]:
from ..payment.models import async_fetch_openrouter_models
models_data = await async_fetch_openrouter_models()
return models_data
DASHBOARD_CSS: str = """
* { margin: 0; padding: 0; box-sizing: border-box; }
body { font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', sans-serif; background: #f5f7fa; color: #2c3e50; line-height: 1.6; padding: 2rem; }
+11 -5
View File
@@ -55,6 +55,9 @@ class ApiKey(SQLModel, table=True): # type: ignore
class ModelRow(SQLModel, table=True): # type: ignore
__tablename__ = "models"
id: str = Field(primary_key=True)
upstream_provider_id: int = Field(
primary_key=True, foreign_key="upstream_providers.id", ondelete="CASCADE"
)
name: str = Field()
created: int = Field()
description: str = Field()
@@ -65,9 +68,6 @@ class ModelRow(SQLModel, table=True): # type: ignore
per_request_limits: str | None = Field(default=None)
top_provider: str | None = Field(default=None)
enabled: bool = Field(default=True, description="Whether this model is enabled")
upstream_provider_id: int | None = Field(
default=None, foreign_key="upstream_providers.id"
)
upstream_provider: "UpstreamProviderRow" = Relationship(back_populates="models")
@@ -75,7 +75,7 @@ class UpstreamProviderRow(SQLModel, table=True): # type: ignore
__tablename__ = "upstream_providers"
id: int | None = Field(default=None, primary_key=True)
provider_type: str = Field(
description="Provider type: generic, openai, azure, openrouter"
description="Provider type: custom, openai, azure, openrouter"
)
base_url: str = Field(unique=True, description="Base URL of the upstream API")
api_key: str = Field(description="API key for the upstream provider")
@@ -83,7 +83,13 @@ class UpstreamProviderRow(SQLModel, table=True): # type: ignore
default=None, description="API version for Azure OpenAI"
)
enabled: bool = Field(default=True, description="Whether this provider is enabled")
models: list["ModelRow"] = Relationship(back_populates="upstream_provider")
provider_fee: float = Field(
default=1.01, description="Provider fee multiplier (default 1%)"
)
models: list["ModelRow"] = Relationship(
back_populates="upstream_provider",
sa_relationship_kwargs={"cascade": "all, delete-orphan"},
)
async def balances_for_mint_and_unit(
+11 -1
View File
@@ -14,6 +14,7 @@ from ..payment.models import (
models_router,
update_sats_pricing,
)
from ..payment.price import update_prices_periodically
from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically
from ..wallet import periodic_payout
from .admin import admin_router
@@ -35,6 +36,7 @@ __version__ = "0.2.0-dev"
async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
logger.info("Application startup initiated", extra={"version": __version__})
btc_price_task = None
pricing_task = None
payout_task = None
nip91_task = None
@@ -65,11 +67,15 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
pass
# await ensure_models_bootstrapped()
await initialize_upstreams()
from ..payment.price import _update_prices
from ..proxy import get_upstreams
from ..upstream import refresh_upstreams_models_periodically
await _update_prices()
await initialize_upstreams()
btc_price_task = asyncio.create_task(update_prices_periodically())
pricing_task = asyncio.create_task(update_sats_pricing())
if global_settings.models_refresh_interval_seconds > 0:
models_refresh_task = asyncio.create_task(
@@ -91,6 +97,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
finally:
logger.info("Application shutdown initiated")
if btc_price_task is not None:
btc_price_task.cancel()
if pricing_task is not None:
pricing_task.cancel()
if payout_task is not None:
@@ -106,6 +114,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
try:
tasks_to_wait = []
if btc_price_task is not None:
tasks_to_wait.append(btc_price_task)
if pricing_task is not None:
tasks_to_wait.append(pricing_task)
if payout_task is not None:
+5 -2
View File
@@ -3,6 +3,7 @@ import math
from pydantic.v1 import BaseModel
from ..core import get_logger
from ..core.db import AsyncSession
from ..core.settings import settings
logger = get_logger(__name__)
@@ -25,7 +26,7 @@ class CostDataError(BaseModel):
async def calculate_cost(
response_data: dict, max_cost: int, session: object | None = None
response_data: dict, max_cost: int, session: AsyncSession
) -> CostData | MaxCostData | CostDataError:
"""
Calculate the cost of an API request based on token usage.
@@ -81,7 +82,9 @@ async def calculate_cost(
from ..upstream import get_model_with_override
upstreams = get_upstreams()
model_obj = await get_model_with_override(response_model, upstreams)
model_obj = await get_model_with_override(
response_model, upstreams, session=session
)
if not model_obj:
logger.error(
+6 -8
View File
@@ -84,7 +84,7 @@ def check_token_balance(headers: dict, body: dict, max_cost_for_model: int) -> N
async def get_max_cost_for_model(
model: str,
session: AsyncSession | None = None,
session: AsyncSession,
model_obj: Any | None = None,
) -> int:
"""Get the maximum cost for a specific model from providers with overrides."""
@@ -109,7 +109,7 @@ async def get_max_cost_for_model(
from ..upstream import get_model_with_override
upstreams = get_upstreams()
model_obj = await get_model_with_override(model, upstreams)
model_obj = await get_model_with_override(model, upstreams, session)
if not model_obj:
fallback_msats = settings.fixed_cost_per_request * 1000
@@ -152,10 +152,10 @@ async def get_max_cost_for_model(
async def calculate_discounted_max_cost(
max_cost_for_model: int, body: dict, session: AsyncSession | None = None
max_cost_for_model: int, body: dict, session: AsyncSession
) -> int:
"""Calculate the discounted max cost for a request using model pricing when available."""
if settings.fixed_pricing or session is None:
if settings.fixed_pricing:
return max_cost_for_model
model = body.get("model", "unknown")
@@ -215,9 +215,7 @@ def estimate_tokens(messages: list) -> int:
return len(str(messages)) // 3
async def get_model_cost_info(
model_id: str, session: AsyncSession | None = None
) -> Pricing | None:
async def get_model_cost_info(model_id: str, session: AsyncSession) -> Pricing | None:
"""Get model pricing info from providers with database overrides."""
if not model_id or model_id == "unknown":
return None
@@ -226,7 +224,7 @@ async def get_model_cost_info(
from ..upstream import get_model_with_override
upstreams = get_upstreams()
model_obj = await get_model_with_override(model_id, upstreams)
model_obj = await get_model_with_override(model_id, upstreams, session)
if model_obj and model_obj.sats_pricing:
return model_obj.sats_pricing
+122 -182
View File
@@ -13,7 +13,7 @@ from sqlmodel.ext.asyncio.session import AsyncSession
from ..core.db import ModelRow, create_session, get_session
from ..core.logging import get_logger
from ..core.settings import settings
from .price import sats_usd_ask_price
from .price import sats_usd_price
logger = get_logger(__name__)
@@ -201,46 +201,56 @@ def load_models() -> list[Model]:
return [Model(**model) for model in models_data] # type: ignore
def _row_to_model(row: ModelRow) -> Model:
def _row_to_model(
row: ModelRow, apply_provider_fee: bool = False, provider_fee: float = 1.01
) -> Model:
architecture = json.loads(row.architecture)
pricing = json.loads(row.pricing)
sats_pricing = json.loads(row.sats_pricing) if row.sats_pricing else None
per_request_limits = (
json.loads(row.per_request_limits) if row.per_request_limits else None
)
top_provider = json.loads(row.top_provider) if row.top_provider else None
top_provider_dict = json.loads(row.top_provider) if row.top_provider else None
# Enforce minimum per-request fee on free/zero-priced models in API output
try:
if isinstance(pricing, dict):
if float(pricing.get("request", 0.0)) <= 0.0:
pricing["request"] = max(pricing.get("request", 0.0), 0.0)
if isinstance(sats_pricing, dict):
if float(sats_pricing.get("request", 0.0)) <= 0.0:
# Convert min_request_msat to sats for sats_pricing fields that are in sats
sats_min = max(1, int(settings.min_request_msat)) / 1000.0
sats_pricing["request"] = max(
sats_pricing.get("request", 0.0), sats_min
)
except Exception:
pass
if apply_provider_fee and isinstance(pricing, dict):
pricing = {k: float(v) * provider_fee for k, v in pricing.items()}
return Model(
if isinstance(pricing, dict) and float(pricing.get("request", 0.0)) <= 0.0:
pricing["request"] = max(pricing.get("request", 0.0), 0.0)
parsed_pricing = Pricing.parse_obj(pricing)
model = Model(
id=row.id,
name=row.name,
created=row.created,
description=row.description,
context_length=row.context_length,
architecture=Architecture.parse_obj(architecture),
pricing=Pricing.parse_obj(pricing),
sats_pricing=Pricing.parse_obj(sats_pricing) if sats_pricing else None,
pricing=parsed_pricing,
sats_pricing=None,
per_request_limits=per_request_limits,
top_provider=TopProvider.parse_obj(top_provider) if top_provider else None,
top_provider=TopProvider.parse_obj(top_provider_dict)
if top_provider_dict
else None,
enabled=row.enabled,
upstream_provider_id=row.upstream_provider_id,
canonical_slug=getattr(row, "canonical_slug", None),
)
if apply_provider_fee:
(
parsed_pricing.max_prompt_cost,
parsed_pricing.max_completion_cost,
parsed_pricing.max_cost,
) = _calculate_usd_max_costs(model)
try:
sats_to_usd = sats_usd_price()
model = _update_model_sats_pricing(model, sats_to_usd)
except Exception as e:
logger.warning(f"Could not calculate sats pricing: {e}")
return model
def _model_to_row_payload(model: Model) -> dict[str, str | int | bool | None]:
return {
@@ -266,33 +276,100 @@ def _model_to_row_payload(model: Model) -> dict[str, str | int | bool | None]:
async def list_models(
session: AsyncSession | None = None,
upstream_id: int | None = None,
session: AsyncSession,
upstream_id: int,
include_disabled: bool = False,
) -> list[Model]:
from sqlmodel import select
from ..core.db import UpstreamProviderRow
query = select(ModelRow)
if upstream_id is not None:
query = query.where(ModelRow.upstream_provider_id == upstream_id)
if not include_disabled:
query = query.where(ModelRow.enabled)
if session is not None:
return [_row_to_model(r) for r in (await session.exec(query)).all()] # type: ignore
async with create_session() as s:
return [_row_to_model(r) for r in (await s.exec(query)).all()] # type: ignore
rows = (await session.exec(query)).all() # type: ignore
provider_result = await session.exec(select(UpstreamProviderRow))
providers_by_id = {p.id: p for p in provider_result.all()}
return [
_row_to_model(
r,
apply_provider_fee=True,
provider_fee=providers_by_id[r.upstream_provider_id].provider_fee
if r.upstream_provider_id in providers_by_id
else 1.01,
)
for r in rows
]
async def get_model_by_id(
model_id: str, session: AsyncSession | None = None
model_id: str, provider_id: int, session: AsyncSession
) -> Model | None:
if session is not None:
row = await session.get(ModelRow, model_id)
return _row_to_model(row) if row and row.enabled else None
async with create_session() as s:
row = await s.get(ModelRow, model_id)
return _row_to_model(row) if row and row.enabled else None
from ..core.db import UpstreamProviderRow
row = await session.get(ModelRow, (model_id, provider_id))
if not row or not row.enabled:
return None
provider = await session.get(UpstreamProviderRow, provider_id)
provider_fee = provider.provider_fee if provider else 1.01
return _row_to_model(row, apply_provider_fee=True, provider_fee=provider_fee)
def _calculate_usd_max_costs(model: Model) -> tuple[float, float, float]:
"""Calculate max costs in USD based on model context/token limits.
Args:
model: Model object
Returns:
Tuple of (max_prompt_cost, max_completion_cost, max_cost) in USD
"""
min_req_msat = max(1, int(getattr(settings, "min_request_msat", 1)))
min_req_usd = float(min_req_msat) / 1_000_000.0
prompt_price = model.pricing.prompt
completion_price = model.pricing.completion
if model.top_provider and (
model.top_provider.context_length or model.top_provider.max_completion_tokens
):
if (cl := model.top_provider.context_length) and (
mct := model.top_provider.max_completion_tokens
):
return (
(cl - mct) * prompt_price,
mct * completion_price,
(cl - mct) * prompt_price + mct * completion_price,
)
elif cl := model.top_provider.context_length:
return (
cl * 0.8 * prompt_price,
cl * 0.2 * completion_price,
cl * prompt_price,
)
elif mct := model.top_provider.max_completion_tokens:
return (
mct * 4 * prompt_price,
mct * completion_price,
mct * 5 * prompt_price,
)
elif model.context_length:
return (
model.context_length * 0.8 * prompt_price,
model.context_length * 0.2 * completion_price,
model.context_length * prompt_price,
)
p = prompt_price * 1_000_000
c = completion_price * 32_000
r = model.pricing.request * 100_000
i = model.pricing.image * 100
w = model.pricing.web_search * 1000
ir = model.pricing.internal_reasoning * 100
return (p, c, max(p + c + r + i + w + ir, min_req_usd))
def _update_model_sats_pricing(model: Model, sats_to_usd: float) -> Model:
@@ -306,59 +383,15 @@ def _update_model_sats_pricing(model: Model, sats_to_usd: float) -> Model:
Updated Model object with new sats_pricing
"""
try:
min_req_msat = max(1, int(getattr(settings, "min_request_msat", 1)))
min_req_sats = float(min_req_msat) / 1000.0
sats = Pricing.parse_obj(
{k: v / sats_to_usd for k, v in model.pricing.dict().items()}
)
min_req_msat = max(1, int(getattr(settings, "min_request_msat", 1)))
min_req_sats = float(min_req_msat) / 1000.0
if sats.request <= 0.0:
sats.request = min_req_sats
mspp = sats.prompt
mspc = sats.completion
if model.top_provider and (
model.top_provider.context_length
or model.top_provider.max_completion_tokens
):
if (cl := model.top_provider.context_length) and (
mct := model.top_provider.max_completion_tokens
):
max_prompt_cost = (cl - mct) * mspp
max_completion_cost = mct * mspc
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
elif cl := model.top_provider.context_length:
max_prompt_cost = cl * 0.8 * mspp
max_completion_cost = cl * 0.2 * mspc
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
elif mct := model.top_provider.max_completion_tokens:
max_prompt_cost = mct * 4 * mspp
max_completion_cost = mct * mspc
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
elif model.context_length:
max_prompt_cost = mspp * model.context_length * 0.8
max_completion_cost = mspc * model.context_length * 0.2
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
else:
p = mspp * 1_000_000
c = mspc * 32_000
r = sats.request * 100_000
i = sats.image * 100
w = sats.web_search * 1000
ir = sats.internal_reasoning * 100
sats.max_prompt_cost = p
sats.max_completion_cost = c
sats.max_cost = p + c + r + i + w + ir
if (sats.max_cost or 0.0) < min_req_sats:
sats.max_cost = min_req_sats
@@ -373,6 +406,9 @@ def _update_model_sats_pricing(model: Model, sats_to_usd: float) -> Model:
sats_pricing=sats,
per_request_limits=model.per_request_limits,
top_provider=model.top_provider,
enabled=model.enabled,
upstream_provider_id=model.upstream_provider_id,
canonical_slug=model.canonical_slug,
)
except Exception as e:
logger.error(
@@ -438,14 +474,13 @@ async def ensure_models_bootstrapped() -> None:
async def _update_sats_pricing_once() -> None:
"""Update sats pricing once for all provider models and database overrides."""
"""Update sats pricing once for all provider models (in-memory only)."""
from ..proxy import get_upstreams
sats_to_usd = await sats_usd_ask_price()
upstreams = get_upstreams()
sats_to_usd = sats_usd_price()
updated_count = 0
for upstream in upstreams:
updated_models = [
_update_model_sats_pricing(m, sats_to_usd)
@@ -455,103 +490,8 @@ async def _update_sats_pricing_once() -> None:
upstream._models_by_id = {m.id: m for m in updated_models}
updated_count += len(updated_models)
async with create_session() as s:
result = await s.exec(
select(ModelRow).where(ModelRow.upstream_provider_id.isnot(None)) # type: ignore
) # type: ignore
rows = result.all()
changed = 0
for row in rows:
try:
pricing = Pricing.parse_obj(json.loads(row.pricing))
top_provider = (
TopProvider.parse_obj(json.loads(row.top_provider))
if row.top_provider
else None
)
sats = Pricing.parse_obj(
{k: v / sats_to_usd for k, v in pricing.dict().items()}
)
min_req_msat = max(1, int(getattr(settings, "min_request_msat", 1)))
min_req_sats = float(min_req_msat) / 1000.0
if sats.request <= 0.0:
sats.request = min_req_sats
mspp = sats.prompt
mspc = sats.completion
if top_provider and (
top_provider.context_length or top_provider.max_completion_tokens
):
if (cl := top_provider.context_length) and (
mct := top_provider.max_completion_tokens
):
max_prompt_cost = (cl - mct) * mspp
max_completion_cost = mct * mspc
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
elif cl := top_provider.context_length:
max_prompt_cost = cl * 0.8 * mspp
max_completion_cost = cl * 0.2 * mspc
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
elif mct := top_provider.max_completion_tokens:
max_prompt_cost = mct * 4 * mspp
max_completion_cost = mct * mspc
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
else:
max_prompt_cost = 1_000_000 * mspp
max_completion_cost = 32_000 * mspc
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
elif row.context_length:
max_prompt_cost = mspp * row.context_length * 0.8
max_completion_cost = mspc * row.context_length * 0.2
sats.max_prompt_cost = max_prompt_cost
sats.max_completion_cost = max_completion_cost
sats.max_cost = max_prompt_cost + max_completion_cost
else:
p = mspp * 1_000_000
c = mspc * 32_000
r = sats.request * 100_000
i = sats.image * 100
w = sats.web_search * 1000
ir = sats.internal_reasoning * 100
sats.max_prompt_cost = p
sats.max_completion_cost = c
sats.max_cost = p + c + r + i + w + ir
if (sats.max_cost or 0.0) < min_req_sats:
sats.max_cost = min_req_sats
new_json = json.dumps(sats.dict())
if row.sats_pricing != new_json:
row.sats_pricing = new_json
s.add(row)
changed += 1
except Exception as per_row_error:
logger.error(
"Failed to update pricing for model",
extra={
"model_id": row.id,
"error": str(per_row_error),
"error_type": type(per_row_error).__name__,
},
)
if changed:
await s.commit()
if updated_count > 0 or changed > 0:
logger.info(
"Updated sats pricing",
extra={
"provider_models_updated": updated_count,
"database_overrides_updated": changed,
},
)
if updated_count > 0:
logger.info("Updated sats pricing", extra={"models_updated": updated_count})
async def update_sats_pricing() -> None:
+64 -31
View File
@@ -1,4 +1,5 @@
import asyncio
import random
import httpx
@@ -7,12 +8,11 @@ from ..core.settings import settings
logger = get_logger(__name__)
def _fees() -> tuple[float, float]:
return settings.exchange_fee, settings.upstream_provider_fee
BTC_USD_PRICE: float | None = None
SATS_USD_PRICE: float | None = None
async def kraken_btc_usd(client: httpx.AsyncClient) -> float | None:
async def _kraken_btc_usd(client: httpx.AsyncClient) -> float | None:
"""Fetch BTC/USD price from Kraken API."""
api = "https://api.kraken.com/0/public/Ticker?pair=XBTUSD"
try:
@@ -33,7 +33,7 @@ async def kraken_btc_usd(client: httpx.AsyncClient) -> float | None:
return None
async def coinbase_btc_usd(client: httpx.AsyncClient) -> float | None:
async def _coinbase_btc_usd(client: httpx.AsyncClient) -> float | None:
"""Fetch BTC/USD price from Coinbase API."""
api = "https://api.coinbase.com/v2/prices/BTC-USD/spot"
try:
@@ -54,7 +54,7 @@ async def coinbase_btc_usd(client: httpx.AsyncClient) -> float | None:
return None
async def binance_btc_usdt(client: httpx.AsyncClient) -> float | None:
async def _binance_btc_usdt(client: httpx.AsyncClient) -> float | None:
"""Fetch BTC/USDT price from Binance API."""
api = "https://api.binance.com/api/v3/ticker/price?symbol=BTCUSDT"
try:
@@ -75,28 +75,20 @@ async def binance_btc_usdt(client: httpx.AsyncClient) -> float | None:
return None
async def btc_usd_ask_price() -> float:
"""Get the lowest BTC/USD price from multiple exchanges with fee adjustment."""
async def _fetch_btc_usd_price() -> float:
"""Fetch the lowest BTC/USD price from multiple exchanges."""
async with httpx.AsyncClient(timeout=30.0) as client:
try:
prices = await asyncio.gather(
kraken_btc_usd(client),
coinbase_btc_usd(client),
binance_btc_usdt(client),
_kraken_btc_usd(client),
_coinbase_btc_usd(client),
_binance_btc_usdt(client),
)
valid_prices = [price for price in prices if price is not None]
if not valid_prices:
logger.error("No valid BTC prices obtained from any exchange")
raise ValueError("Unable to fetch BTC price from any exchange")
min_price = min(valid_prices)
exchange_fee, provider_fee = _fees()
final_price = min_price / (exchange_fee * provider_fee)
return final_price
return min(valid_prices)
except Exception as e:
logger.error(
"Error in BTC price aggregation",
@@ -105,18 +97,59 @@ async def btc_usd_ask_price() -> float:
raise
async def sats_usd_ask_price() -> float:
"""Get the USD price per satoshi."""
async def _update_prices() -> None:
"""Update global BTC and SATS price variables."""
global BTC_USD_PRICE, SATS_USD_PRICE
btc_price = await _fetch_btc_usd_price()
BTC_USD_PRICE = btc_price
SATS_USD_PRICE = btc_price / 100_000_000
logger.info(
"Updated BTC/USD price",
extra={"btc_usd": btc_price, "sats_usd": SATS_USD_PRICE},
)
def btc_usd_price() -> float:
"""Get the current BTC/USD price."""
if BTC_USD_PRICE is None:
raise ValueError("BTC price not initialized")
return BTC_USD_PRICE
def sats_usd_price() -> float:
"""Get the current USD price per satoshi."""
if SATS_USD_PRICE is None:
raise ValueError("SATS price not initialized")
return SATS_USD_PRICE
async def update_prices_periodically() -> None:
"""Background task to periodically update BTC and SATS prices."""
try:
btc_price = await btc_usd_ask_price()
sats_price = btc_price / 100_000_000
if not settings.enable_pricing_refresh:
return
except Exception:
pass
return sats_price
await _update_prices()
except Exception as e:
logger.error(
"Error calculating satoshi price",
extra={"error": str(e), "error_type": type(e).__name__},
)
raise
while True:
try:
interval = getattr(settings, "pricing_refresh_interval_seconds", 120)
jitter = max(0.0, float(interval) * 0.1)
await asyncio.sleep(interval + random.uniform(0, jitter))
except asyncio.CancelledError:
break
try:
if not settings.enable_pricing_refresh:
return
except Exception:
pass
try:
await _update_prices()
except asyncio.CancelledError:
break
except Exception as e:
logger.error(f"Error updating BTC/SATS prices: {e}")
+47 -19
View File
@@ -7,7 +7,14 @@ from sqlmodel import select
from .auth import pay_for_request, revert_pay_for_request, validate_bearer_key
from .core import get_logger
from .core.db import ApiKey, AsyncSession, ModelRow, create_session, get_session
from .core.db import (
ApiKey,
AsyncSession,
ModelRow,
UpstreamProviderRow,
create_session,
get_session,
)
from .payment.helpers import (
calculate_discounted_max_cost,
check_token_balance,
@@ -82,8 +89,19 @@ async def refresh_model_maps() -> None:
async with create_session() as session:
result = await session.exec(select(ModelRow).where(ModelRow.enabled))
override_rows = result.all()
overrides_by_id = {
row.id: row for row in override_rows if row.upstream_provider_id is not None
provider_result = await session.exec(select(UpstreamProviderRow))
providers_by_id = {p.id: p for p in provider_result.all()}
overrides_by_id: dict[str, tuple[ModelRow, float]] = {
row.id: (
row,
providers_by_id[row.upstream_provider_id].provider_fee
if row.upstream_provider_id in providers_by_id
else 1.01,
)
for row in override_rows
if row.upstream_provider_id is not None
}
for upstream in _upstreams:
@@ -99,15 +117,20 @@ async def refresh_model_maps() -> None:
if openrouter:
for model in openrouter.get_cached_models():
if model.enabled:
model_to_use = (
_row_to_model(overrides_by_id[model.id])
if model.id in overrides_by_id
else model
)
if model.id in overrides_by_id:
override_row, provider_fee = overrides_by_id[model.id]
model_to_use = _row_to_model(
override_row, apply_provider_fee=True, provider_fee=provider_fee
)
else:
model_to_use = model
base_id = get_base_model_id(model_to_use.id)
if base_id not in unique_models:
unique_models[base_id] = model_to_use
for alias in resolve_model_alias(model.id, model_to_use.canonical_slug):
unique_model = model_to_use.copy(update={"id": base_id})
unique_models[base_id] = unique_model
for alias in resolve_model_alias(
model_to_use.id, model_to_use.canonical_slug
):
model_instances[alias] = model_to_use
provider_map[alias] = openrouter
@@ -115,18 +138,23 @@ async def refresh_model_maps() -> None:
upstream_prefix = getattr(upstream, "upstream_name", None)
for model in upstream.get_cached_models():
if model.enabled:
model_to_use = (
_row_to_model(overrides_by_id[model.id])
if model.id in overrides_by_id
else model
)
if model.id in overrides_by_id:
override_row, provider_fee = overrides_by_id[model.id]
model_to_use = _row_to_model(
override_row, apply_provider_fee=True, provider_fee=provider_fee
)
else:
model_to_use = model
base_id = get_base_model_id(model_to_use.id)
unique_models[base_id] = model_to_use
unique_model = model_to_use.copy(update={"id": base_id})
unique_models[base_id] = unique_model
aliases = resolve_model_alias(model.id, model_to_use.canonical_slug)
aliases = resolve_model_alias(
model_to_use.id, model_to_use.canonical_slug
)
if upstream_prefix and "/" not in model.id:
prefixed_id = f"{upstream_prefix}/{model.id}"
if upstream_prefix and "/" not in model_to_use.id:
prefixed_id = f"{upstream_prefix}/{model_to_use.id}"
if prefixed_id not in aliases:
aliases.append(prefixed_id)
+144 -31
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
import json
import os
import re
import traceback
from collections.abc import AsyncGenerator
@@ -90,8 +91,19 @@ async def get_all_models_with_overrides(
async with create_session() as session:
result = await session.exec(select(ModelRow).where(ModelRow.enabled))
override_rows = result.all()
overrides_by_id = {
row.id: row for row in override_rows if row.upstream_provider_id is not None
provider_result = await session.exec(select(UpstreamProviderRow))
providers_by_id = {p.id: p for p in provider_result.all()}
overrides_by_id: dict[str, tuple[ModelRow, float]] = {
row.id: (
row,
providers_by_id[row.upstream_provider_id].provider_fee
if row.upstream_provider_id in providers_by_id
else 1.01,
)
for row in override_rows
if row.upstream_provider_id is not None
}
all_models: dict[str, Model] = {}
@@ -99,7 +111,10 @@ async def get_all_models_with_overrides(
for upstream in upstreams:
for model in upstream.get_cached_models():
if model.id in overrides_by_id:
all_models[model.id] = _row_to_model(overrides_by_id[model.id])
override_row, provider_fee = overrides_by_id[model.id]
all_models[model.id] = _row_to_model(
override_row, apply_provider_fee=True, provider_fee=provider_fee
)
elif model.enabled:
all_models[model.id] = model
@@ -109,6 +124,7 @@ async def get_all_models_with_overrides(
async def get_model_with_override(
model_id: str,
upstreams: list[UpstreamProvider],
session: AsyncSession,
) -> Model | None:
"""Get a specific model from providers with database override applied.
@@ -127,18 +143,23 @@ async def get_model_with_override(
aliases = resolve_model_alias(model_id)
async with create_session() as session:
for alias in aliases:
result = await session.exec(
select(ModelRow).where(
ModelRow.id == alias,
ModelRow.upstream_provider_id.isnot(None), # type: ignore
ModelRow.enabled,
)
for alias in aliases:
result = await session.exec(
select(ModelRow).where(
ModelRow.id == alias,
ModelRow.upstream_provider_id.isnot(None), # type: ignore
ModelRow.enabled,
)
)
override_row = result.first()
if override_row:
provider = await session.get(
UpstreamProviderRow, override_row.upstream_provider_id
)
provider_fee = provider.provider_fee if provider else 1.01
return _row_to_model(
override_row, apply_provider_fee=True, provider_fee=provider_fee
)
override_row = result.first()
if override_row:
return _row_to_model(override_row)
for alias in aliases:
for upstream in upstreams:
@@ -192,9 +213,6 @@ async def refresh_upstreams_models_periodically(
break
import os
async def init_upstreams() -> list[UpstreamProvider]:
"""Initialize upstream providers from database.
@@ -355,7 +373,7 @@ async def _seed_providers_from_settings(
else:
providers_to_add.append(
UpstreamProviderRow(
provider_type="generic",
provider_type="custom",
base_url=base_url,
api_key=settings.upstream_api_key,
enabled=True,
@@ -382,7 +400,9 @@ def _instantiate_provider(provider_row: UpstreamProviderRow) -> UpstreamProvider
"""
try:
if provider_row.provider_type == "openai":
return OpenAIUpstreamProvider(provider_row.api_key)
return OpenAIUpstreamProvider(
provider_row.api_key, provider_row.provider_fee
)
elif provider_row.provider_type == "azure":
if not provider_row.api_version:
logger.error(
@@ -394,11 +414,16 @@ def _instantiate_provider(provider_row: UpstreamProviderRow) -> UpstreamProvider
provider_row.base_url,
provider_row.api_key,
provider_row.api_version,
provider_row.provider_fee,
)
elif provider_row.provider_type == "openrouter":
return OpenRouterUpstreamProvider(provider_row.api_key)
elif provider_row.provider_type == "generic":
return UpstreamProvider(provider_row.base_url, provider_row.api_key)
return OpenRouterUpstreamProvider(
provider_row.api_key, provider_row.provider_fee
)
elif provider_row.provider_type == "custom":
return UpstreamProvider(
provider_row.base_url, provider_row.api_key, provider_row.provider_fee
)
else:
logger.error(
f"Unknown provider type: {provider_row.provider_type}",
@@ -423,18 +448,21 @@ class UpstreamProvider:
base_url: str
api_key: str
upstream_name: str | None = None
provider_fee: float = 1.05
_models_cache: list[Model] = []
_models_by_id: dict[str, Model] = {}
def __init__(self, base_url: str, api_key: str):
def __init__(self, base_url: str, api_key: str, provider_fee: float = 1.01):
"""Initialize the upstream provider.
Args:
base_url: Base URL of the upstream API endpoint
api_key: API key for authenticating with the upstream service
provider_fee: Provider fee multiplier (default 1.01 for 1% fee)
"""
self.base_url = base_url
self.api_key = api_key
self.provider_fee = provider_fee
self._models_cache = []
self._models_by_id = {}
@@ -1973,6 +2001,59 @@ class UpstreamProvider:
token=x_cashu_token,
)
def _apply_provider_fee_to_model(self, model: Model) -> Model:
"""Apply provider fee to model's USD pricing and calculate max costs.
Args:
model: Model object to update
Returns:
Model with provider fee applied to pricing and max costs calculated
"""
from .payment.models import Pricing, _calculate_usd_max_costs
adjusted_pricing = Pricing.parse_obj(
{k: v * self.provider_fee for k, v in model.pricing.dict().items()}
)
temp_model = Model(
id=model.id,
name=model.name,
created=model.created,
description=model.description,
context_length=model.context_length,
architecture=model.architecture,
pricing=adjusted_pricing,
sats_pricing=None,
per_request_limits=model.per_request_limits,
top_provider=model.top_provider,
enabled=model.enabled,
upstream_provider_id=model.upstream_provider_id,
canonical_slug=model.canonical_slug,
)
(
adjusted_pricing.max_prompt_cost,
adjusted_pricing.max_completion_cost,
adjusted_pricing.max_cost,
) = _calculate_usd_max_costs(temp_model)
return Model(
id=model.id,
name=model.name,
created=model.created,
description=model.description,
context_length=model.context_length,
architecture=model.architecture,
pricing=adjusted_pricing,
sats_pricing=model.sats_pricing,
per_request_limits=model.per_request_limits,
top_provider=model.top_provider,
enabled=model.enabled,
upstream_provider_id=model.upstream_provider_id,
canonical_slug=model.canonical_slug,
)
async def fetch_models(self) -> list[Model]:
"""Fetch available models from upstream API and update cache.
@@ -1985,9 +2066,21 @@ class UpstreamProvider:
async def refresh_models_cache(self) -> None:
"""Refresh the in-memory models cache from upstream API."""
try:
from .payment.models import _update_model_sats_pricing
from .payment.price import sats_usd_price
models = await self.fetch_models()
self._models_cache = models
self._models_by_id = {m.id: m for m in models}
models_with_fees = [self._apply_provider_fee_to_model(m) for m in models]
try:
sats_to_usd = sats_usd_price()
self._models_cache = [
_update_model_sats_pricing(m, sats_to_usd) for m in models_with_fees
]
except Exception:
self._models_cache = models_with_fees
self._models_by_id = {m.id: m for m in self._models_cache}
logger.info(
f"Refreshed models cache for {self.upstream_name or self.base_url}",
extra={"model_count": len(models)},
@@ -2021,9 +2114,13 @@ class UpstreamProvider:
class OpenAIUpstreamProvider(UpstreamProvider):
"""Upstream provider specifically configured for OpenAI API."""
def __init__(self, api_key: str):
def __init__(self, api_key: str, provider_fee: float = 1.01):
self.upstream_name = "openai"
super().__init__(base_url="https://api.openai.com/v1", api_key=api_key)
super().__init__(
base_url="https://api.openai.com/v1",
api_key=api_key,
provider_fee=provider_fee,
)
def transform_model_name(self, model_id: str) -> str:
"""Strip 'openai/' prefix for OpenAI API compatibility."""
@@ -2038,15 +2135,26 @@ class OpenAIUpstreamProvider(UpstreamProvider):
class AzureUpstreamProvider(UpstreamProvider):
"""Upstream provider specifically configured for Azure OpenAI Service."""
def __init__(self, base_url: str, api_key: str, api_version: str):
def __init__(
self,
base_url: str,
api_key: str,
api_version: str,
provider_fee: float = 1.01,
):
"""Initialize Azure provider with API key and version.
Args:
base_url: Azure OpenAI endpoint base URL
api_key: Azure OpenAI API key for authentication
api_version: Azure OpenAI API version (e.g., "2024-02-15-preview")
provider_fee: Provider fee multiplier (default 1.01 for 1% fee)
"""
super().__init__(base_url=base_url, api_key=api_key)
super().__init__(
base_url=base_url,
api_key=api_key,
provider_fee=provider_fee,
)
self.api_version = api_version
def prepare_params(
@@ -2070,14 +2178,19 @@ class AzureUpstreamProvider(UpstreamProvider):
class OpenRouterUpstreamProvider(UpstreamProvider):
"""Upstream provider specifically configured for OpenRouter API."""
def __init__(self, api_key: str):
def __init__(self, api_key: str, provider_fee: float = 1.06):
"""Initialize OpenRouter provider with API key.
Args:
api_key: OpenRouter API key for authentication
provider_fee: Provider fee multiplier (default 1.06 for 6% fee)
"""
self.upstream_name = "openrouter"
super().__init__(base_url="https://openrouter.ai/api/v1", api_key=api_key)
super().__init__(
base_url="https://openrouter.ai/api/v1",
api_key=api_key,
provider_fee=provider_fee,
)
async def fetch_models(self) -> list[Model]:
"""Fetch all OpenRouter models."""