open-webui/backend/open_webui/models/functions.py
Timothy Jaeryang Baek e6476928ee refac
2026-10-09 18:37:55 +04:00

476 lines
19 KiB
Python

"""Function (filter/action/pipe) models, forms, and database operations."""
from __future__ import annotations
import logging
import time
# local imports
from fastapi import HTTPException
from open_webui.internal.db import Base, JSONField, get_async_db_context
from open_webui.models.function_history import FunctionHistories, FunctionHistory, function_snapshot
from open_webui.models.users import User, UserResponse, Users, UserSettings
from open_webui.utils.valves import decrypt_valves, encrypt_valves, validate_valves
from pydantic import BaseModel, ConfigDict
from sqlalchemy import BigInteger, Boolean, Column, Index, String, Text, delete, select, update
from sqlalchemy.ext.asyncio import AsyncSession
log = logging.getLogger(__name__)
class Function(Base): # database table mapping
__tablename__ = 'function'
id = Column(String, primary_key=True, unique=True)
version_id = Column(Text, nullable=True)
user_id = Column(String, index=True) # creator user id
name = Column(Text, nullable=False) # function identifier
type = Column(Text, nullable=False) # function type (pipe, filter, etc.)
content = Column(Text, nullable=True) # Python source code
meta = Column(JSONField, nullable=True) # function metadata
valves = Column(JSONField, nullable=True) # function configuration valves
is_active = Column(Boolean, default=False) # function activation status
is_global = Column(Boolean) # if True, applied to every chat automatically
updated_at = Column(BigInteger) # epoch seconds
created_at = Column(BigInteger) # epoch seconds
__table_args__ = (Index('is_global_idx', 'is_global'),) # speed up global-function lookups
class FunctionMeta(BaseModel):
description: str | None = None
manifest: dict | None = {}
model_config = ConfigDict(extra='allow')
class FunctionModel(BaseModel):
version_id: str | None = None
id: str
user_id: str | None = None # may be null for legacy/malformed records
name: str
type: str
content: str
meta: FunctionMeta
is_active: bool = False
is_global: bool = False
updated_at: int # timestamp in epoch
created_at: int # timestamp in epoch
model_config = ConfigDict(from_attributes=True) # allows ORM model binding
# --- form / schema definitions ---
class FunctionWithValvesModel(BaseModel):
version_id: str | None = None
id: str
user_id: str | None = None # may be null for legacy/malformed records
name: str
type: str
content: str
meta: FunctionMeta
valves: dict | None = None
is_active: bool = False
is_global: bool = False
updated_at: int # timestamp in epoch
created_at: int # timestamp in epoch
model_config = ConfigDict(from_attributes=True)
####################
# Forms
####################
class FunctionResponse(BaseModel):
version_id: str | None = None
id: str
user_id: str | None = None # may be null for legacy/malformed records
type: str
name: str
meta: FunctionMeta
is_active: bool
is_global: bool
updated_at: int # timestamp in epoch
created_at: int # timestamp in epoch
model_config = ConfigDict(from_attributes=True)
class FunctionUserResponse(FunctionResponse):
user: UserResponse | None = None
class FunctionForm(BaseModel):
commit_message: str | None = None
id: str
name: str
content: str
meta: FunctionMeta
class FunctionValves(BaseModel):
valves: dict | None = None
class FunctionsTable:
async def _lock_function(self, session, id):
# UPDATE also serializes writers on SQLite, where SELECT FOR UPDATE does not.
await session.execute(update(Function).where(Function.id == id).values(version_id=Function.version_id))
return await session.get(Function, id, populate_existing=True)
async def _write_function(
self,
session,
resource,
updated,
user_id=None,
version_id=None,
module=None,
merge_meta=False,
):
updated = dict(updated)
message = updated.pop('commit_message', None)
updated.pop('version_id', None) # Imported pointers never belong to this resource.
before = function_snapshot(resource)
if version_id:
entry = (
await session.execute(select(FunctionHistory).filter_by(id=version_id, function_id=resource.id))
).scalar_one_or_none()
if not entry:
raise HTTPException(404, 'Version not found')
# The prepared candidate must be the exact selected saved configuration.
if function_snapshot(updated) != function_snapshot(entry.snapshot):
raise HTTPException(400, 'Version configuration does not match the saved snapshot')
if module is not None:
validate_valves(module, updated.get('valves', resource.valves))
if merge_meta:
updated['meta'] = {**(resource.meta or {}), **updated.get('meta', {})}
for key, value in updated.items():
setattr(resource, key, value)
after = function_snapshot(resource)
if version_id:
resource.version_id = version_id
elif after != before or not resource.version_id:
entry = FunctionHistories.new_entry(
resource.id, after, user_id or resource.user_id or '', resource.version_id, message
)
session.add(entry)
resource.version_id = entry.id
resource.updated_at = int(time.time())
async def insert_new_function(self, user_id, type, form_data, db=None, module=None):
async with get_async_db_context(db) as session:
try:
function = Function(
**form_data.model_dump(exclude={'commit_message'}),
user_id=user_id,
type=type,
is_active=False,
is_global=False,
updated_at=int(time.time()),
created_at=int(time.time()),
)
session.add(function)
await self._write_function(
session,
function,
{'commit_message': form_data.commit_message},
user_id,
module=module,
)
await session.flush()
result = FunctionModel.model_validate(function)
await session.commit()
return result
except Exception:
await session.rollback()
raise
async def sync_functions(self, user_id, functions, db=None, modules=None):
async with get_async_db_context(db) as session:
try:
# Lock all existing rows in a stable order before applying the batch.
ids = (await session.execute(select(Function.id).order_by(Function.id))).scalars().all()
existing = {id: await self._lock_function(session, id) for id in ids}
incoming = {func.id for func in functions}
for func in functions:
data = func.model_dump(exclude={'version_id'})
data['valves'] = encrypt_valves(data.get('valves'))
data['user_id'] = user_id
resource = existing.get(func.id)
if resource is None:
resource = Function(**data)
session.add(resource)
await self._write_function(
session,
resource,
data,
user_id,
module=(modules or {}).get(func.id),
)
for id in set(existing) - incoming:
await session.execute(delete(FunctionHistory).filter_by(function_id=id))
await session.delete(existing[id])
await session.flush()
rows = (await session.execute(select(Function))).scalars().all()
result = [
FunctionWithValvesModel.model_validate(
{**FunctionModel.model_validate(row).model_dump(), 'valves': decrypt_valves(row.valves)}
)
for row in rows
]
await session.commit()
return result
except Exception:
await session.rollback()
raise
async def get_function_by_id(self, id: str, db: AsyncSession | None = None) -> FunctionModel | None:
try:
async with get_async_db_context(db) as db:
function = await db.get(Function, id)
return FunctionModel.model_validate(function) if function else None
except Exception:
return None
async def get_functions_by_ids(self, ids: list[str], db: AsyncSession | None = None) -> list[FunctionModel]:
"""
Batch fetch multiple functions by their IDs in a single query.
Returns functions in the same order as the input IDs (None entries filtered out).
"""
if not ids:
return []
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Function).filter(Function.id.in_(ids)))
functions = result.scalars().all()
# Create a dict for O(1) lookup
func_dict = {f.id: FunctionModel.model_validate(f) for f in functions}
# Return in original order, filtering out any not found
return [func_dict[id] for id in ids if id in func_dict]
except Exception:
return []
async def get_functions(
self, active_only=False, include_valves=False, db: AsyncSession | None = None
) -> list[FunctionModel | FunctionWithValvesModel]:
async with get_async_db_context(db) as db:
if active_only:
result = await db.execute(select(Function).filter_by(is_active=True))
else:
result = await db.execute(select(Function))
functions = result.scalars().all()
if include_valves:
return [
FunctionWithValvesModel.model_validate(
{
**FunctionModel.model_validate(function).model_dump(),
'valves': decrypt_valves(function.valves),
}
)
for function in functions
]
else:
return [FunctionModel.model_validate(function) for function in functions]
async def get_function_list(self, db: AsyncSession | None = None) -> list[FunctionUserResponse]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Function).order_by(Function.updated_at.desc()))
functions = result.scalars().all()
user_ids = list(set(func.user_id for func in functions))
users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
users_dict = {user.id: user for user in users}
return [
FunctionUserResponse.model_validate(
{
**FunctionResponse.model_validate(func).model_dump(),
'user': (
UserResponse(
id=users_dict[func.user_id].id,
name=users_dict[func.user_id].name,
role=users_dict[func.user_id].role,
email=users_dict[func.user_id].email,
).model_dump()
if func.user_id in users_dict
else None
),
}
)
for func in functions
]
async def get_functions_by_type(
self, type: str, active_only=False, db: AsyncSession | None = None
) -> list[FunctionModel]:
async with get_async_db_context(db) as db:
if active_only:
result = await db.execute(select(Function).filter_by(type=type, is_active=True))
else:
result = await db.execute(select(Function).filter_by(type=type))
return [FunctionModel.model_validate(function) for function in result.scalars().all()]
async def get_global_filter_functions(self, db: AsyncSession | None = None) -> list[FunctionModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Function).filter_by(type='filter', is_active=True, is_global=True))
return [FunctionModel.model_validate(function) for function in result.scalars().all()]
async def get_active_function_ids_by_type(
self, type: str, db: AsyncSession | None = None
) -> list[tuple[str, bool]]:
"""Return (id, is_global) for active functions without fetching plugin source."""
async with get_async_db_context(db) as db:
result = await db.execute(select(Function.id, Function.is_global).filter_by(type=type, is_active=True))
return [(id, bool(is_global)) for id, is_global in result.all()]
async def get_active_filter_ids(self, db: AsyncSession | None = None) -> list[tuple[str, bool]]:
"""Return (id, is_global) for active filters without fetching plugin source."""
return await self.get_active_function_ids_by_type('filter', db=db)
async def get_global_action_functions(self, db: AsyncSession | None = None) -> list[FunctionModel]:
async with get_async_db_context(db) as db:
result = await db.execute(select(Function).filter_by(type='action', is_active=True, is_global=True))
return [FunctionModel.model_validate(function) for function in result.scalars().all()]
async def get_function_valves_by_id(self, id: str, db: AsyncSession | None = None) -> dict | None:
async with get_async_db_context(db) as db:
try:
result = await db.execute(select(Function.valves).filter_by(id=id))
return decrypt_valves(result.scalar_one_or_none())
except Exception as e:
log.exception(f'Error getting function valves by id {id}: {e}')
return None
async def get_function_valves_by_ids(self, ids: list[str], db: AsyncSession | None = None) -> dict[str, dict]:
"""
Batch fetch valves for multiple functions in a single query.
Returns a dict mapping function_id -> valves dict.
Functions without valves are mapped to {}.
"""
if not ids:
return {}
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(Function.id, Function.valves).filter(Function.id.in_(ids)))
return {id: decrypt_valves(valves) for id, valves in result.all()}
except Exception as e:
log.exception(f'Error batch-fetching function valves: {e}')
return {}
async def update_function_valves_by_id(
self, id: str, valves: dict, db: AsyncSession | None = None
) -> FunctionValves | None:
async with get_async_db_context(db) as db:
try:
function = await db.get(Function, id)
function.valves = encrypt_valves(valves)
function.updated_at = int(time.time())
await db.commit()
return FunctionModel.model_validate(function)
except Exception:
return None
async def update_function_metadata_by_id(self, id, metadata, db=None, user_id=None):
return await self.update_function_by_id(id, {'meta': metadata}, db=db, user_id=user_id, merge_meta=True)
async def get_user_valves_by_id_and_user_id(
self, id: str, user_id: str, db: AsyncSession | None = None
) -> dict | None:
try:
async with get_async_db_context(db) as db:
result = await db.execute(select(User.settings).filter_by(id=user_id))
settings = result.scalar_one_or_none()
user_settings = UserSettings(**settings).model_dump() if settings else {}
# Check if user has "functions" and "valves" settings
if 'functions' not in user_settings:
user_settings['functions'] = {}
if 'valves' not in user_settings['functions']:
user_settings['functions']['valves'] = {}
return decrypt_valves(user_settings['functions']['valves'].get(id))
except Exception:
log.exception(f'Error getting user values by id {id} and user id {user_id}')
return None
async def update_user_valves_by_id_and_user_id(
self, id: str, user_id: str, valves: dict, db: AsyncSession | None = None
) -> dict | None:
try:
user = await Users.get_user_by_id(user_id, db=db)
user_settings = user.settings.model_dump() if user.settings else {}
# Check if user has "functions" and "valves" settings
if 'functions' not in user_settings:
user_settings['functions'] = {}
if 'valves' not in user_settings['functions']:
user_settings['functions']['valves'] = {}
user_settings['functions']['valves'][id] = encrypt_valves(valves)
# Update the user settings in the database
await Users.update_user_by_id(user_id, {'settings': user_settings}, db=db)
return valves
except Exception as e:
log.exception(f'Error updating user valves by id {id} and user_id {user_id}: {e}')
return None
async def update_function_by_id(
self, id, updated, db=None, user_id=None, version_id=None, module=None, merge_meta=False
):
async with get_async_db_context(db) as session:
try:
function = await self._lock_function(session, id)
if not function:
raise ValueError('Function not found')
await self._write_function(
session,
function,
updated,
user_id,
version_id,
module,
merge_meta,
)
await session.flush()
result = FunctionModel.model_validate(function)
await session.commit()
return result
except Exception:
await session.rollback()
raise
async def deactivate_all_functions(self, db: AsyncSession | None = None) -> bool | None:
async with get_async_db_context(db) as db:
try:
await db.execute(
update(Function).values(
is_active=False,
updated_at=int(time.time()),
)
)
await db.commit()
return True
except Exception:
return None
async def delete_function_by_id(self, id: str, db: AsyncSession | None = None) -> bool:
async with get_async_db_context(db) as db:
try:
await self._lock_function(db, id)
await db.execute(delete(FunctionHistory).filter_by(function_id=id))
await db.execute(delete(Function).filter_by(id=id))
await db.commit()
return True
except Exception:
return False
Functions = FunctionsTable() # singleton functions engine