mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-08-02 00:26:13 +00:00
refactor pricing, provider fees, realtime model map updates
This commit is contained in:
@@ -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
@@ -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()">×</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
@@ -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
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user