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> <script>
let providersList = []; let providersList = [];
let selectedProviderId = null; let selectedProviderId = null;
let providerModels = { db_models: [], remote_models: [] }; let providerModels = { db_models: [], remote_models: [], provider: {} };
let openrouterPresets = [];
async function fetchProviders() { async function fetchProviders() {
const tableBody = document.getElementById('providers-tbody'); 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 { try {
const resp = await fetch('/admin/api/upstream-providers', { credentials: 'same-origin' }); const resp = await fetch('/admin/api/upstream-providers', { credentials: 'same-origin' });
if (!resp.ok) throw new Error('HTTP ' + resp.status); if (!resp.ok) throw new Error('HTTP ' + resp.status);
providersList = await resp.json(); providersList = await resp.json();
renderProvidersTable(); renderProvidersTable();
} catch (e) { } 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() { function renderProvidersTable() {
const tableBody = document.getElementById('providers-tbody'); const tableBody = document.getElementById('providers-tbody');
if (!Array.isArray(providersList) || !providersList.length) { 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; return;
} }
const rows = providersList.map(p => ` const rows = providersList.map(p => `
<tr> <tr>
<td>${p.id}</td>
<td>${p.provider_type}</td> <td>${p.provider_type}</td>
<td style="word-break: break-all;">${p.base_url}</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> <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() { function renderProviderModels() {
const dbModelsBody = document.getElementById('db-models-tbody'); const dbModelsBody = document.getElementById('db-models-tbody');
const remoteModelsBody = document.getElementById('remote-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) { if (providerModels.db_models && providerModels.db_models.length > 0) {
dbModelsBody.innerHTML = ''; dbModelsBody.innerHTML = '';
@@ -1483,32 +1486,40 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
dbModelsBody.appendChild(row); dbModelsBody.appendChild(row);
}); });
} else { } 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) { if (isCustomProvider) {
remoteModelsBody.innerHTML = ''; remoteModelsSection.style.display = 'none';
providerModels.remote_models.forEach(m => { customProviderActions.style.display = 'block';
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 { } 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 }, 'openrouter': { baseUrl: 'https://openrouter.ai/api/v1', showApiVersion: false },
'anthropic': { baseUrl: 'https://api.anthropic.com/v1', showApiVersion: false }, 'anthropic': { baseUrl: 'https://api.anthropic.com/v1', showApiVersion: false },
'azure': { baseUrl: '', showApiVersion: true }, '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() { function updateProviderFields() {
const providerType = document.getElementById('provider-type').value; const providerType = document.getElementById('provider-type').value;
const config = PROVIDER_CONFIGS[providerType] || { baseUrl: '', showApiVersion: false }; const config = PROVIDER_CONFIGS[providerType] || { baseUrl: '', showApiVersion: false };
const baseUrlField = document.getElementById('provider-base-url'); const baseUrlField = document.getElementById('provider-base-url');
const apiVersionRow = document.getElementById('api-version-row'); const apiVersionRow = document.getElementById('api-version-row');
const providerId = document.getElementById('provider-id').value; const providerId = document.getElementById('provider-id').value;
const feeField = document.getElementById('provider-fee');
if (config.baseUrl && !providerId) { if (config.baseUrl && !providerId) {
baseUrlField.value = config.baseUrl; baseUrlField.value = config.baseUrl;
@@ -1542,6 +1558,10 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
} }
apiVersionRow.style.display = config.showApiVersion ? 'block' : 'none'; apiVersionRow.style.display = config.showApiVersion ? 'block' : 'none';
if (!feeField.value) {
feeField.placeholder = getProviderFeePlaceholder(providerType);
}
} }
async function openProviderEditor(providerId) { 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-key').placeholder = '[Keep existing]';
document.getElementById('provider-api-version').value = p.api_version || ''; document.getElementById('provider-api-version').value = p.api_version || '';
document.getElementById('provider-enabled').checked = p.enabled; document.getElementById('provider-enabled').checked = p.enabled;
document.getElementById('provider-fee').value = p.provider_fee || '';
updateProviderFields(); updateProviderFields();
} catch (e) { } catch (e) {
errorBox.style.display = 'block'; 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-key').placeholder = 'API Key';
document.getElementById('provider-api-version').value = ''; document.getElementById('provider-api-version').value = '';
document.getElementById('provider-enabled').checked = true; document.getElementById('provider-enabled').checked = true;
document.getElementById('provider-fee').value = '';
document.getElementById('provider-fee').placeholder = getProviderFeePlaceholder('openrouter');
updateProviderFields(); updateProviderFields();
} }
@@ -1596,11 +1619,16 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
const errorBox = document.getElementById('provider-error'); const errorBox = document.getElementById('provider-error');
errorBox.style.display = 'none'; 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 = { const payload = {
provider_type: document.getElementById('provider-type').value, provider_type: providerType,
base_url: document.getElementById('provider-base-url').value, base_url: document.getElementById('provider-base-url').value,
api_version: document.getElementById('provider-api-version').value || null, api_version: document.getElementById('provider-api-version').value || null,
enabled: document.getElementById('provider-enabled').checked, enabled: document.getElementById('provider-enabled').checked,
provider_fee: feeValue ? parseFloat(feeValue) : defaultFee,
}; };
const apiKey = document.getElementById('provider-api-key').value; 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) { if (enabled) {
const modal = document.getElementById('model-override-modal'); const modal = document.getElementById('model-override-modal');
const errorBox = document.getElementById('model-override-error'); const errorBox = document.getElementById('model-override-error');
errorBox.style.display = 'none'; errorBox.style.display = 'none';
document.getElementById('override-model-id').value = modelData.id || ''; 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-model-name').value = modelData.name || '';
document.getElementById('override-description').value = modelData.description || ''; document.getElementById('override-description').value = modelData.description || '';
document.getElementById('override-context').value = modelData.context_length || 8192; document.getElementById('override-context').value = modelData.context_length || 8192;
@@ -1699,9 +1842,10 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
prompt: 0.0, prompt: 0.0,
completion: 0.0, completion: 0.0,
request: 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_prompt_cost;
delete pricing.max_completion_cost; delete pricing.max_completion_cost;
delete pricing.max_cost; delete pricing.max_cost;
@@ -1711,8 +1855,9 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
document.getElementById('override-mode').value = 'create'; document.getElementById('override-mode').value = 'create';
document.getElementById('override-created').value = Math.floor(Date.now() / 1000); document.getElementById('override-created').value = Math.floor(Date.now() / 1000);
document.getElementById('override-upstream-provider-id').value = selectedProviderId; document.getElementById('override-upstream-provider-id').value = selectedProviderId;
document.getElementById('modal-title').textContent = 'Create Model Override'; const isCustomProvider = providerModels.provider && providerModels.provider.provider_type === 'custom';
document.getElementById('override-save-btn').textContent = 'Create Override'; 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; document.getElementById('override-model-name').disabled = false;
modal.style.display = 'block'; modal.style.display = 'block';
@@ -1735,7 +1880,7 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
enabled: false enabled: false
}; };
const resp = await fetch('/admin/api/models', { const resp = await fetch(`/admin/api/upstream-providers/${selectedProviderId}/models`, {
method: 'POST', method: 'POST',
headers: { 'Content-Type': 'application/json' }, headers: { 'Content-Type': 'application/json' },
credentials: 'same-origin', credentials: 'same-origin',
@@ -1756,7 +1901,7 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
async function toggleModelEnabled(modelId, newEnabledState) { async function toggleModelEnabled(modelId, newEnabledState) {
try { 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' credentials: 'same-origin'
}); });
if (!resp.ok) throw new Error('Failed to fetch model'); if (!resp.ok) throw new Error('Failed to fetch model');
@@ -1764,7 +1909,7 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
model.enabled = newEnabledState; 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', method: 'PATCH',
headers: { 'Content-Type': 'application/json' }, headers: { 'Content-Type': 'application/json' },
credentials: 'same-origin', credentials: 'same-origin',
@@ -1809,7 +1954,7 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
}; };
const resp = await fetch( 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', method: isEdit ? 'PATCH' : 'POST',
headers: { 'Content-Type': 'application/json' }, headers: { 'Content-Type': 'application/json' },
@@ -1841,7 +1986,7 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
errorBox.style.display = 'none'; errorBox.style.display = 'none';
try { 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' credentials: 'same-origin'
}); });
if (!resp.ok) throw new Error('Failed to fetch model'); 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; if (!confirm('Delete this model override?')) return;
try { try {
const resp = await fetch(`/admin/api/models/${encodeURIComponent(modelId)}`, { const resp = await fetch(`/admin/api/upstream-providers/${selectedProviderId}/models/${encodeURIComponent(modelId)}`, {
method: 'DELETE', method: 'DELETE',
credentials: 'same-origin' credentials: 'same-origin'
}); });
@@ -1907,8 +2052,10 @@ UPSTREAM_PROVIDERS_JS: str = """<!--html-->
window.onclick = function(event) { window.onclick = function(event) {
const editModal = document.getElementById('provider-edit-modal'); const editModal = document.getElementById('provider-edit-modal');
const overrideModal = document.getElementById('model-override-modal'); const overrideModal = document.getElementById('model-override-modal');
const presetModal = document.getElementById('preset-selector-modal');
if (event.target == editModal) closeProviderEditor(); if (event.target == editModal) closeProviderEditor();
if (event.target == overrideModal) closeModelOverrideModal(); if (event.target == overrideModal) closeModelOverrideModal();
if (event.target == presetModal) closePresetSelector();
} }
</script> </script>
""" """
@@ -1936,7 +2083,6 @@ def upstream_providers_page() -> str:
<table> <table>
<thead> <thead>
<tr> <tr>
<th>ID</th>
<th>Type</th> <th>Type</th>
<th>Base URL</th> <th>Base URL</th>
<th>Status</th> <th>Status</th>
@@ -1944,7 +2090,7 @@ def upstream_providers_page() -> str:
</tr> </tr>
</thead> </thead>
<tbody id="providers-tbody"> <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> </tbody>
</table> </table>
</div> </div>
@@ -1958,7 +2104,15 @@ def upstream_providers_page() -> str:
<div id="models-loading" style="color:#718096;">Loading models…</div> <div id="models-loading" style="color:#718096;">Loading models…</div>
<div id="models-content" style="display:none;"> <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;"> <table style="margin-bottom: 2rem;">
<thead> <thead>
<tr> <tr>
@@ -1968,23 +2122,25 @@ def upstream_providers_page() -> str:
</tr> </tr>
</thead> </thead>
<tbody id="db-models-tbody"> <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> </tbody>
</table> </table>
<h3>Remote Models</h3> <div id="remote-models-section">
<table> <h3>Remote Models</h3>
<thead> <table>
<tr> <thead>
<th>Model ID</th> <tr>
<th>Name</th> <th>Model ID</th>
<th>Actions</th> <th>Name</th>
</tr> <th>Actions</th>
</thead> </tr>
<tbody id="remote-models-tbody"> </thead>
<tr><td colspan="3" style="color:#718096;">No remote models</td></tr> <tbody id="remote-models-tbody">
</tbody> <tr><td colspan="3" style="color:#718096;">No remote models</td></tr>
</table> </tbody>
</table>
</div>
</div> </div>
</div> </div>
@@ -2003,7 +2159,7 @@ def upstream_providers_page() -> str:
<option value="openai">OpenAI</option> <option value="openai">OpenAI</option>
<option value="anthropic">Anthropic</option> <option value="anthropic">Anthropic</option>
<option value="azure">Azure OpenAI</option> <option value="azure">Azure OpenAI</option>
<option value="generic">Generic</option> <option value="custom">Custom</option>
</select> </select>
<label>Base URL</label> <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"> <input type="text" id="provider-api-version" placeholder="2024-02-15-preview">
</div> </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;"> <label style="display:flex; align-items:center; gap:8px; margin:10px 0;">
<input type="checkbox" id="provider-enabled" style="width:auto;"> <input type="checkbox" id="provider-enabled" style="width:auto;">
<span>Enabled</span> <span>Enabled</span>
@@ -2065,6 +2225,48 @@ def upstream_providers_page() -> str:
</div> </div>
</div> </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> </body>
</html> </html>
""" """
@@ -2078,12 +2280,6 @@ async def admin_upstream_providers(request: Request) -> str:
return admin_auth() 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): class ModelCreate(BaseModel):
id: str id: str
name: str name: str
@@ -2098,14 +2294,25 @@ class ModelCreate(BaseModel):
enabled: bool = True enabled: bool = True
@admin_router.post("/api/models", dependencies=[Depends(require_admin_api)]) @admin_router.post(
async def create_model_admin_api(payload: ModelCreate) -> dict[str, object]: "/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: 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: if exists:
raise HTTPException( 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( row = ModelRow(
id=payload.id, id=payload.id,
name=payload.name, name=payload.name,
@@ -2123,7 +2330,7 @@ async def create_model_admin_api(payload: ModelCreate) -> dict[str, object]:
top_provider=( top_provider=(
json.dumps(payload.top_provider) if payload.top_provider else None 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, enabled=payload.enabled,
) )
session.add(row) session.add(row)
@@ -2131,70 +2338,29 @@ async def create_model_admin_api(payload: ModelCreate) -> dict[str, object]:
await session.refresh(row) await session.refresh(row)
await refresh_model_maps() 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.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}
@admin_router.get( @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: 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: if not row:
raise HTTPException(status_code=404, detail="Model not found") raise HTTPException(
return _row_to_model(row).dict() # type: ignore 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): class ModelUpdate(BaseModel):
@@ -2212,18 +2378,25 @@ class ModelUpdate(BaseModel):
@admin_router.patch( @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( async def update_provider_model(
model_id: str, payload: ModelUpdate provider_id: int, model_id: str, payload: ModelUpdate
) -> dict[str, object]: ) -> dict[str, object]:
if payload.id != model_id: if payload.id != model_id:
raise HTTPException(status_code=400, detail="Path id does not match payload id") raise HTTPException(status_code=400, detail="Path id does not match payload id")
async with create_session() as session: 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: 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.name = payload.name
row.description = payload.description row.description = payload.description
@@ -2240,7 +2413,6 @@ async def update_model_admin_api(
row.top_provider = ( row.top_provider = (
json.dumps(payload.top_provider) if payload.top_provider else None json.dumps(payload.top_provider) if payload.top_provider else None
) )
row.upstream_provider_id = payload.upstream_provider_id
row.enabled = payload.enabled row.enabled = payload.enabled
session.add(row) session.add(row)
@@ -2248,33 +2420,43 @@ async def update_model_admin_api(
await session.refresh(row) await session.refresh(row)
await refresh_model_maps() 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( @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: 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: 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.delete(row)
await session.commit() await session.commit()
await refresh_model_maps() await refresh_model_maps()
return {"ok": True, "deleted_id": model_id} return {"ok": True, "deleted_id": model_id}
@admin_router.delete("/api/models", dependencies=[Depends(require_admin_api)]) @admin_router.delete(
async def delete_all_models_admin_api() -> dict[str, object]: "/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: 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() rows = result.all()
for row in rows: for row in rows:
await session.delete(row) # type: ignore await session.delete(row) # type: ignore
await session.commit() await session.commit()
await refresh_model_maps() await refresh_model_maps()
return {"ok": True, "deleted": "all"} return {"ok": True, "deleted": len(rows)}
class UpstreamProviderCreate(BaseModel): class UpstreamProviderCreate(BaseModel):
@@ -2283,6 +2465,7 @@ class UpstreamProviderCreate(BaseModel):
api_key: str api_key: str
api_version: str | None = None api_version: str | None = None
enabled: bool = True enabled: bool = True
provider_fee: float = 1.01
class UpstreamProviderUpdate(BaseModel): class UpstreamProviderUpdate(BaseModel):
@@ -2291,6 +2474,7 @@ class UpstreamProviderUpdate(BaseModel):
api_key: str | None = None api_key: str | None = None
api_version: str | None = None api_version: str | None = None
enabled: bool | None = None enabled: bool | None = None
provider_fee: float | None = None
@admin_router.get("/api/upstream-providers", dependencies=[Depends(require_admin_api)]) @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_key": "[REDACTED]" if p.api_key else "",
"api_version": p.api_version, "api_version": p.api_version,
"enabled": p.enabled, "enabled": p.enabled,
"provider_fee": p.provider_fee,
} }
for p in providers for p in providers
] ]
@@ -2332,12 +2517,14 @@ async def create_upstream_provider(
api_key=payload.api_key, api_key=payload.api_key,
api_version=payload.api_version, api_version=payload.api_version,
enabled=payload.enabled, enabled=payload.enabled,
provider_fee=payload.provider_fee,
) )
session.add(provider) session.add(provider)
await session.commit() await session.commit()
await session.refresh(provider) await session.refresh(provider)
await reinitialize_upstreams() await reinitialize_upstreams()
await refresh_model_maps()
return { return {
"id": provider.id, "id": provider.id,
"provider_type": provider.provider_type, "provider_type": provider.provider_type,
@@ -2345,6 +2532,7 @@ async def create_upstream_provider(
"api_key": "[REDACTED]", "api_key": "[REDACTED]",
"api_version": provider.api_version, "api_version": provider.api_version,
"enabled": provider.enabled, "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_key": "[REDACTED]" if provider.api_key else "",
"api_version": provider.api_version, "api_version": provider.api_version,
"enabled": provider.enabled, "enabled": provider.enabled,
"provider_fee": provider.provider_fee,
} }
@@ -2387,12 +2576,15 @@ async def update_upstream_provider(
provider.api_version = payload.api_version provider.api_version = payload.api_version
if payload.enabled is not None: if payload.enabled is not None:
provider.enabled = payload.enabled provider.enabled = payload.enabled
if payload.provider_fee is not None:
provider.provider_fee = payload.provider_fee
session.add(provider) session.add(provider)
await session.commit() await session.commit()
await session.refresh(provider) await session.refresh(provider)
await reinitialize_upstreams() await reinitialize_upstreams()
await refresh_model_maps()
return { return {
"id": provider.id, "id": provider.id,
"provider_type": provider.provider_type, "provider_type": provider.provider_type,
@@ -2400,6 +2592,7 @@ async def update_upstream_provider(
"api_key": "[REDACTED]", "api_key": "[REDACTED]",
"api_version": provider.api_version, "api_version": provider.api_version,
"enabled": provider.enabled, "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.delete(provider)
await session.commit() await session.commit()
await reinitialize_upstreams() await reinitialize_upstreams()
await refresh_model_maps()
return {"ok": True, "deleted_id": provider_id} 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) upstream_instance = _instantiate_provider(provider)
if upstream_instance: if upstream_instance:
try: try:
models = await upstream_instance.fetch_models() raw_models = await upstream_instance.fetch_models()
remote_models = [m.dict() for m in models] remote_models = [
upstream_instance._apply_provider_fee_to_model(m).dict()
for m in raw_models
]
except Exception as e: except Exception as e:
logger.error( logger.error(
f"Failed to fetch models from {provider.provider_type}: {e}" 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 = """ DASHBOARD_CSS: str = """
* { margin: 0; padding: 0; box-sizing: border-box; } * { 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; } 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 class ModelRow(SQLModel, table=True): # type: ignore
__tablename__ = "models" __tablename__ = "models"
id: str = Field(primary_key=True) id: str = Field(primary_key=True)
upstream_provider_id: int = Field(
primary_key=True, foreign_key="upstream_providers.id", ondelete="CASCADE"
)
name: str = Field() name: str = Field()
created: int = Field() created: int = Field()
description: str = Field() description: str = Field()
@@ -65,9 +68,6 @@ class ModelRow(SQLModel, table=True): # type: ignore
per_request_limits: str | None = Field(default=None) per_request_limits: str | None = Field(default=None)
top_provider: str | None = Field(default=None) top_provider: str | None = Field(default=None)
enabled: bool = Field(default=True, description="Whether this model is enabled") 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") upstream_provider: "UpstreamProviderRow" = Relationship(back_populates="models")
@@ -75,7 +75,7 @@ class UpstreamProviderRow(SQLModel, table=True): # type: ignore
__tablename__ = "upstream_providers" __tablename__ = "upstream_providers"
id: int | None = Field(default=None, primary_key=True) id: int | None = Field(default=None, primary_key=True)
provider_type: str = Field( 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") base_url: str = Field(unique=True, description="Base URL of the upstream API")
api_key: str = Field(description="API key for the upstream provider") 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" default=None, description="API version for Azure OpenAI"
) )
enabled: bool = Field(default=True, description="Whether this provider is enabled") 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( async def balances_for_mint_and_unit(
+11 -1
View File
@@ -14,6 +14,7 @@ from ..payment.models import (
models_router, models_router,
update_sats_pricing, update_sats_pricing,
) )
from ..payment.price import update_prices_periodically
from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically from ..proxy import initialize_upstreams, proxy_router, refresh_model_maps_periodically
from ..wallet import periodic_payout from ..wallet import periodic_payout
from .admin import admin_router from .admin import admin_router
@@ -35,6 +36,7 @@ __version__ = "0.2.0-dev"
async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]: async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
logger.info("Application startup initiated", extra={"version": __version__}) logger.info("Application startup initiated", extra={"version": __version__})
btc_price_task = None
pricing_task = None pricing_task = None
payout_task = None payout_task = None
nip91_task = None nip91_task = None
@@ -65,11 +67,15 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
pass pass
# await ensure_models_bootstrapped() # await ensure_models_bootstrapped()
await initialize_upstreams()
from ..payment.price import _update_prices
from ..proxy import get_upstreams from ..proxy import get_upstreams
from ..upstream import refresh_upstreams_models_periodically 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()) pricing_task = asyncio.create_task(update_sats_pricing())
if global_settings.models_refresh_interval_seconds > 0: if global_settings.models_refresh_interval_seconds > 0:
models_refresh_task = asyncio.create_task( models_refresh_task = asyncio.create_task(
@@ -91,6 +97,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
finally: finally:
logger.info("Application shutdown initiated") logger.info("Application shutdown initiated")
if btc_price_task is not None:
btc_price_task.cancel()
if pricing_task is not None: if pricing_task is not None:
pricing_task.cancel() pricing_task.cancel()
if payout_task is not None: if payout_task is not None:
@@ -106,6 +114,8 @@ async def lifespan(_: FastAPI) -> AsyncGenerator[None, None]:
try: try:
tasks_to_wait = [] tasks_to_wait = []
if btc_price_task is not None:
tasks_to_wait.append(btc_price_task)
if pricing_task is not None: if pricing_task is not None:
tasks_to_wait.append(pricing_task) tasks_to_wait.append(pricing_task)
if payout_task is not None: if payout_task is not None:
+5 -2
View File
@@ -3,6 +3,7 @@ import math
from pydantic.v1 import BaseModel from pydantic.v1 import BaseModel
from ..core import get_logger from ..core import get_logger
from ..core.db import AsyncSession
from ..core.settings import settings from ..core.settings import settings
logger = get_logger(__name__) logger = get_logger(__name__)
@@ -25,7 +26,7 @@ class CostDataError(BaseModel):
async def calculate_cost( 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: ) -> CostData | MaxCostData | CostDataError:
""" """
Calculate the cost of an API request based on token usage. 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 from ..upstream import get_model_with_override
upstreams = get_upstreams() 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: if not model_obj:
logger.error( 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( async def get_max_cost_for_model(
model: str, model: str,
session: AsyncSession | None = None, session: AsyncSession,
model_obj: Any | None = None, model_obj: Any | None = None,
) -> int: ) -> int:
"""Get the maximum cost for a specific model from providers with overrides.""" """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 from ..upstream import get_model_with_override
upstreams = get_upstreams() 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: if not model_obj:
fallback_msats = settings.fixed_cost_per_request * 1000 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( 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: ) -> int:
"""Calculate the discounted max cost for a request using model pricing when available.""" """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 return max_cost_for_model
model = body.get("model", "unknown") model = body.get("model", "unknown")
@@ -215,9 +215,7 @@ def estimate_tokens(messages: list) -> int:
return len(str(messages)) // 3 return len(str(messages)) // 3
async def get_model_cost_info( async def get_model_cost_info(model_id: str, session: AsyncSession) -> Pricing | None:
model_id: str, session: AsyncSession | None = None
) -> Pricing | None:
"""Get model pricing info from providers with database overrides.""" """Get model pricing info from providers with database overrides."""
if not model_id or model_id == "unknown": if not model_id or model_id == "unknown":
return None return None
@@ -226,7 +224,7 @@ async def get_model_cost_info(
from ..upstream import get_model_with_override from ..upstream import get_model_with_override
upstreams = get_upstreams() 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: if model_obj and model_obj.sats_pricing:
return 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.db import ModelRow, create_session, get_session
from ..core.logging import get_logger from ..core.logging import get_logger
from ..core.settings import settings from ..core.settings import settings
from .price import sats_usd_ask_price from .price import sats_usd_price
logger = get_logger(__name__) logger = get_logger(__name__)
@@ -201,46 +201,56 @@ def load_models() -> list[Model]:
return [Model(**model) for model in models_data] # type: ignore 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) architecture = json.loads(row.architecture)
pricing = json.loads(row.pricing) pricing = json.loads(row.pricing)
sats_pricing = json.loads(row.sats_pricing) if row.sats_pricing else None
per_request_limits = ( per_request_limits = (
json.loads(row.per_request_limits) if row.per_request_limits else None 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 if apply_provider_fee and isinstance(pricing, dict):
try: pricing = {k: float(v) * provider_fee for k, v in pricing.items()}
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
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, id=row.id,
name=row.name, name=row.name,
created=row.created, created=row.created,
description=row.description, description=row.description,
context_length=row.context_length, context_length=row.context_length,
architecture=Architecture.parse_obj(architecture), architecture=Architecture.parse_obj(architecture),
pricing=Pricing.parse_obj(pricing), pricing=parsed_pricing,
sats_pricing=Pricing.parse_obj(sats_pricing) if sats_pricing else None, sats_pricing=None,
per_request_limits=per_request_limits, 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, enabled=row.enabled,
upstream_provider_id=row.upstream_provider_id, upstream_provider_id=row.upstream_provider_id,
canonical_slug=getattr(row, "canonical_slug", None), 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]: def _model_to_row_payload(model: Model) -> dict[str, str | int | bool | None]:
return { return {
@@ -266,33 +276,100 @@ def _model_to_row_payload(model: Model) -> dict[str, str | int | bool | None]:
async def list_models( async def list_models(
session: AsyncSession | None = None, session: AsyncSession,
upstream_id: int | None = None, upstream_id: int,
include_disabled: bool = False, include_disabled: bool = False,
) -> list[Model]: ) -> list[Model]:
from sqlmodel import select from sqlmodel import select
from ..core.db import UpstreamProviderRow
query = select(ModelRow) query = select(ModelRow)
if upstream_id is not None: if upstream_id is not None:
query = query.where(ModelRow.upstream_provider_id == upstream_id) query = query.where(ModelRow.upstream_provider_id == upstream_id)
if not include_disabled: if not include_disabled:
query = query.where(ModelRow.enabled) query = query.where(ModelRow.enabled)
if session is not None: rows = (await session.exec(query)).all() # type: ignore
return [_row_to_model(r) for r in (await session.exec(query)).all()] # type: ignore provider_result = await session.exec(select(UpstreamProviderRow))
async with create_session() as s: providers_by_id = {p.id: p for p in provider_result.all()}
return [_row_to_model(r) for r in (await s.exec(query)).all()] # type: ignore 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( async def get_model_by_id(
model_id: str, session: AsyncSession | None = None model_id: str, provider_id: int, session: AsyncSession
) -> Model | None: ) -> Model | None:
if session is not None: from ..core.db import UpstreamProviderRow
row = await session.get(ModelRow, model_id)
return _row_to_model(row) if row and row.enabled else None row = await session.get(ModelRow, (model_id, provider_id))
async with create_session() as s: if not row or not row.enabled:
row = await s.get(ModelRow, model_id) return None
return _row_to_model(row) if row and row.enabled else 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: 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 Updated Model object with new sats_pricing
""" """
try: 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( sats = Pricing.parse_obj(
{k: v / sats_to_usd for k, v in model.pricing.dict().items()} {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: if sats.request <= 0.0:
sats.request = min_req_sats 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: if (sats.max_cost or 0.0) < min_req_sats:
sats.max_cost = 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, sats_pricing=sats,
per_request_limits=model.per_request_limits, per_request_limits=model.per_request_limits,
top_provider=model.top_provider, top_provider=model.top_provider,
enabled=model.enabled,
upstream_provider_id=model.upstream_provider_id,
canonical_slug=model.canonical_slug,
) )
except Exception as e: except Exception as e:
logger.error( logger.error(
@@ -438,14 +474,13 @@ async def ensure_models_bootstrapped() -> None:
async def _update_sats_pricing_once() -> 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 from ..proxy import get_upstreams
sats_to_usd = await sats_usd_ask_price()
upstreams = get_upstreams() upstreams = get_upstreams()
sats_to_usd = sats_usd_price()
updated_count = 0 updated_count = 0
for upstream in upstreams: for upstream in upstreams:
updated_models = [ updated_models = [
_update_model_sats_pricing(m, sats_to_usd) _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} upstream._models_by_id = {m.id: m for m in updated_models}
updated_count += len(updated_models) updated_count += len(updated_models)
async with create_session() as s: if updated_count > 0:
result = await s.exec( logger.info("Updated sats pricing", extra={"models_updated": updated_count})
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,
},
)
async def update_sats_pricing() -> None: async def update_sats_pricing() -> None:
+64 -31
View File
@@ -1,4 +1,5 @@
import asyncio import asyncio
import random
import httpx import httpx
@@ -7,12 +8,11 @@ from ..core.settings import settings
logger = get_logger(__name__) logger = get_logger(__name__)
BTC_USD_PRICE: float | None = None
def _fees() -> tuple[float, float]: SATS_USD_PRICE: float | None = None
return settings.exchange_fee, settings.upstream_provider_fee
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.""" """Fetch BTC/USD price from Kraken API."""
api = "https://api.kraken.com/0/public/Ticker?pair=XBTUSD" api = "https://api.kraken.com/0/public/Ticker?pair=XBTUSD"
try: try:
@@ -33,7 +33,7 @@ async def kraken_btc_usd(client: httpx.AsyncClient) -> float | None:
return 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.""" """Fetch BTC/USD price from Coinbase API."""
api = "https://api.coinbase.com/v2/prices/BTC-USD/spot" api = "https://api.coinbase.com/v2/prices/BTC-USD/spot"
try: try:
@@ -54,7 +54,7 @@ async def coinbase_btc_usd(client: httpx.AsyncClient) -> float | None:
return 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.""" """Fetch BTC/USDT price from Binance API."""
api = "https://api.binance.com/api/v3/ticker/price?symbol=BTCUSDT" api = "https://api.binance.com/api/v3/ticker/price?symbol=BTCUSDT"
try: try:
@@ -75,28 +75,20 @@ async def binance_btc_usdt(client: httpx.AsyncClient) -> float | None:
return None return None
async def btc_usd_ask_price() -> float: async def _fetch_btc_usd_price() -> float:
"""Get the lowest BTC/USD price from multiple exchanges with fee adjustment.""" """Fetch the lowest BTC/USD price from multiple exchanges."""
async with httpx.AsyncClient(timeout=30.0) as client: async with httpx.AsyncClient(timeout=30.0) as client:
try: try:
prices = await asyncio.gather( prices = await asyncio.gather(
kraken_btc_usd(client), _kraken_btc_usd(client),
coinbase_btc_usd(client), _coinbase_btc_usd(client),
binance_btc_usdt(client), _binance_btc_usdt(client),
) )
valid_prices = [price for price in prices if price is not None] valid_prices = [price for price in prices if price is not None]
if not valid_prices: if not valid_prices:
logger.error("No valid BTC prices obtained from any exchange") logger.error("No valid BTC prices obtained from any exchange")
raise ValueError("Unable to fetch BTC price from any exchange") raise ValueError("Unable to fetch BTC price from any exchange")
return min(valid_prices)
min_price = min(valid_prices)
exchange_fee, provider_fee = _fees()
final_price = min_price / (exchange_fee * provider_fee)
return final_price
except Exception as e: except Exception as e:
logger.error( logger.error(
"Error in BTC price aggregation", "Error in BTC price aggregation",
@@ -105,18 +97,59 @@ async def btc_usd_ask_price() -> float:
raise raise
async def sats_usd_ask_price() -> float: async def _update_prices() -> None:
"""Get the USD price per satoshi.""" """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: try:
btc_price = await btc_usd_ask_price() if not settings.enable_pricing_refresh:
sats_price = btc_price / 100_000_000 return
except Exception:
pass
return sats_price await _update_prices()
except Exception as e: while True:
logger.error( try:
"Error calculating satoshi price", interval = getattr(settings, "pricing_refresh_interval_seconds", 120)
extra={"error": str(e), "error_type": type(e).__name__}, jitter = max(0.0, float(interval) * 0.1)
) await asyncio.sleep(interval + random.uniform(0, jitter))
raise 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 .auth import pay_for_request, revert_pay_for_request, validate_bearer_key
from .core import get_logger 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 ( from .payment.helpers import (
calculate_discounted_max_cost, calculate_discounted_max_cost,
check_token_balance, check_token_balance,
@@ -82,8 +89,19 @@ async def refresh_model_maps() -> None:
async with create_session() as session: async with create_session() as session:
result = await session.exec(select(ModelRow).where(ModelRow.enabled)) result = await session.exec(select(ModelRow).where(ModelRow.enabled))
override_rows = result.all() 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: for upstream in _upstreams:
@@ -99,15 +117,20 @@ async def refresh_model_maps() -> None:
if openrouter: if openrouter:
for model in openrouter.get_cached_models(): for model in openrouter.get_cached_models():
if model.enabled: if model.enabled:
model_to_use = ( if model.id in overrides_by_id:
_row_to_model(overrides_by_id[model.id]) override_row, provider_fee = overrides_by_id[model.id]
if model.id in overrides_by_id model_to_use = _row_to_model(
else 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) base_id = get_base_model_id(model_to_use.id)
if base_id not in unique_models: if base_id not in unique_models:
unique_models[base_id] = model_to_use unique_model = model_to_use.copy(update={"id": base_id})
for alias in resolve_model_alias(model.id, model_to_use.canonical_slug): 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 model_instances[alias] = model_to_use
provider_map[alias] = openrouter provider_map[alias] = openrouter
@@ -115,18 +138,23 @@ async def refresh_model_maps() -> None:
upstream_prefix = getattr(upstream, "upstream_name", None) upstream_prefix = getattr(upstream, "upstream_name", None)
for model in upstream.get_cached_models(): for model in upstream.get_cached_models():
if model.enabled: if model.enabled:
model_to_use = ( if model.id in overrides_by_id:
_row_to_model(overrides_by_id[model.id]) override_row, provider_fee = overrides_by_id[model.id]
if model.id in overrides_by_id model_to_use = _row_to_model(
else 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) 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: if upstream_prefix and "/" not in model_to_use.id:
prefixed_id = f"{upstream_prefix}/{model.id}" prefixed_id = f"{upstream_prefix}/{model_to_use.id}"
if prefixed_id not in aliases: if prefixed_id not in aliases:
aliases.append(prefixed_id) aliases.append(prefixed_id)
+144 -31
View File
@@ -1,6 +1,7 @@
from __future__ import annotations from __future__ import annotations
import json import json
import os
import re import re
import traceback import traceback
from collections.abc import AsyncGenerator from collections.abc import AsyncGenerator
@@ -90,8 +91,19 @@ async def get_all_models_with_overrides(
async with create_session() as session: async with create_session() as session:
result = await session.exec(select(ModelRow).where(ModelRow.enabled)) result = await session.exec(select(ModelRow).where(ModelRow.enabled))
override_rows = result.all() 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] = {} all_models: dict[str, Model] = {}
@@ -99,7 +111,10 @@ async def get_all_models_with_overrides(
for upstream in upstreams: for upstream in upstreams:
for model in upstream.get_cached_models(): for model in upstream.get_cached_models():
if model.id in overrides_by_id: 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: elif model.enabled:
all_models[model.id] = model all_models[model.id] = model
@@ -109,6 +124,7 @@ async def get_all_models_with_overrides(
async def get_model_with_override( async def get_model_with_override(
model_id: str, model_id: str,
upstreams: list[UpstreamProvider], upstreams: list[UpstreamProvider],
session: AsyncSession,
) -> Model | None: ) -> Model | None:
"""Get a specific model from providers with database override applied. """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) aliases = resolve_model_alias(model_id)
async with create_session() as session: for alias in aliases:
for alias in aliases: result = await session.exec(
result = await session.exec( select(ModelRow).where(
select(ModelRow).where( ModelRow.id == alias,
ModelRow.id == alias, ModelRow.upstream_provider_id.isnot(None), # type: ignore
ModelRow.upstream_provider_id.isnot(None), # type: ignore ModelRow.enabled,
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 alias in aliases:
for upstream in upstreams: for upstream in upstreams:
@@ -192,9 +213,6 @@ async def refresh_upstreams_models_periodically(
break break
import os
async def init_upstreams() -> list[UpstreamProvider]: async def init_upstreams() -> list[UpstreamProvider]:
"""Initialize upstream providers from database. """Initialize upstream providers from database.
@@ -355,7 +373,7 @@ async def _seed_providers_from_settings(
else: else:
providers_to_add.append( providers_to_add.append(
UpstreamProviderRow( UpstreamProviderRow(
provider_type="generic", provider_type="custom",
base_url=base_url, base_url=base_url,
api_key=settings.upstream_api_key, api_key=settings.upstream_api_key,
enabled=True, enabled=True,
@@ -382,7 +400,9 @@ def _instantiate_provider(provider_row: UpstreamProviderRow) -> UpstreamProvider
""" """
try: try:
if provider_row.provider_type == "openai": 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": elif provider_row.provider_type == "azure":
if not provider_row.api_version: if not provider_row.api_version:
logger.error( logger.error(
@@ -394,11 +414,16 @@ def _instantiate_provider(provider_row: UpstreamProviderRow) -> UpstreamProvider
provider_row.base_url, provider_row.base_url,
provider_row.api_key, provider_row.api_key,
provider_row.api_version, provider_row.api_version,
provider_row.provider_fee,
) )
elif provider_row.provider_type == "openrouter": elif provider_row.provider_type == "openrouter":
return OpenRouterUpstreamProvider(provider_row.api_key) return OpenRouterUpstreamProvider(
elif provider_row.provider_type == "generic": provider_row.api_key, provider_row.provider_fee
return UpstreamProvider(provider_row.base_url, provider_row.api_key) )
elif provider_row.provider_type == "custom":
return UpstreamProvider(
provider_row.base_url, provider_row.api_key, provider_row.provider_fee
)
else: else:
logger.error( logger.error(
f"Unknown provider type: {provider_row.provider_type}", f"Unknown provider type: {provider_row.provider_type}",
@@ -423,18 +448,21 @@ class UpstreamProvider:
base_url: str base_url: str
api_key: str api_key: str
upstream_name: str | None = None upstream_name: str | None = None
provider_fee: float = 1.05
_models_cache: list[Model] = [] _models_cache: list[Model] = []
_models_by_id: dict[str, 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. """Initialize the upstream provider.
Args: Args:
base_url: Base URL of the upstream API endpoint base_url: Base URL of the upstream API endpoint
api_key: API key for authenticating with the upstream service 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.base_url = base_url
self.api_key = api_key self.api_key = api_key
self.provider_fee = provider_fee
self._models_cache = [] self._models_cache = []
self._models_by_id = {} self._models_by_id = {}
@@ -1973,6 +2001,59 @@ class UpstreamProvider:
token=x_cashu_token, 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]: async def fetch_models(self) -> list[Model]:
"""Fetch available models from upstream API and update cache. """Fetch available models from upstream API and update cache.
@@ -1985,9 +2066,21 @@ class UpstreamProvider:
async def refresh_models_cache(self) -> None: async def refresh_models_cache(self) -> None:
"""Refresh the in-memory models cache from upstream API.""" """Refresh the in-memory models cache from upstream API."""
try: try:
from .payment.models import _update_model_sats_pricing
from .payment.price import sats_usd_price
models = await self.fetch_models() models = await self.fetch_models()
self._models_cache = models models_with_fees = [self._apply_provider_fee_to_model(m) for m in models]
self._models_by_id = {m.id: 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( logger.info(
f"Refreshed models cache for {self.upstream_name or self.base_url}", f"Refreshed models cache for {self.upstream_name or self.base_url}",
extra={"model_count": len(models)}, extra={"model_count": len(models)},
@@ -2021,9 +2114,13 @@ class UpstreamProvider:
class OpenAIUpstreamProvider(UpstreamProvider): class OpenAIUpstreamProvider(UpstreamProvider):
"""Upstream provider specifically configured for OpenAI API.""" """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" 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: def transform_model_name(self, model_id: str) -> str:
"""Strip 'openai/' prefix for OpenAI API compatibility.""" """Strip 'openai/' prefix for OpenAI API compatibility."""
@@ -2038,15 +2135,26 @@ class OpenAIUpstreamProvider(UpstreamProvider):
class AzureUpstreamProvider(UpstreamProvider): class AzureUpstreamProvider(UpstreamProvider):
"""Upstream provider specifically configured for Azure OpenAI Service.""" """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. """Initialize Azure provider with API key and version.
Args: Args:
base_url: Azure OpenAI endpoint base URL base_url: Azure OpenAI endpoint base URL
api_key: Azure OpenAI API key for authentication api_key: Azure OpenAI API key for authentication
api_version: Azure OpenAI API version (e.g., "2024-02-15-preview") 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 self.api_version = api_version
def prepare_params( def prepare_params(
@@ -2070,14 +2178,19 @@ class AzureUpstreamProvider(UpstreamProvider):
class OpenRouterUpstreamProvider(UpstreamProvider): class OpenRouterUpstreamProvider(UpstreamProvider):
"""Upstream provider specifically configured for OpenRouter API.""" """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. """Initialize OpenRouter provider with API key.
Args: Args:
api_key: OpenRouter API key for authentication api_key: OpenRouter API key for authentication
provider_fee: Provider fee multiplier (default 1.06 for 6% fee)
""" """
self.upstream_name = "openrouter" 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]: async def fetch_models(self) -> list[Model]:
"""Fetch all OpenRouter models.""" """Fetch all OpenRouter models."""