Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 22 additions & 12 deletions singlestoredb/apps/_python_udfs.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,17 +38,6 @@ async def run_udf_app(

app_config = AppConfig.from_env()

if kill_existing_app_server:
# Shutdown the server gracefully if it was started by us.
# Since the uvicorn server doesn't start a new subprocess
# killing the process would result in kernel dying.
if _running_server is not None:
await _running_server.shutdown()
_running_server = None

# Kill if any other process is occupying the port
kill_process_by_port(app_config.listen_port)

base_url = generate_base_url(app_config)

udf_suffix = ''
Expand All @@ -68,6 +57,21 @@ async def run_udf_app(
f'You can only define a maximum of {MAX_UDFS_LIMIT} functions.',
)

# Prove every name is allowed before killing a live interactive server.
if app_config.running_interactively:
app.preflight_interactive_functions()

if kill_existing_app_server:
# Shutdown the server gracefully if it was started by us.
# Since the uvicorn server doesn't start a new subprocess
# killing the process would result in kernel dying.
if _running_server is not None:
await _running_server.shutdown()
_running_server = None

# Kill if any other process is occupying the port
kill_process_by_port(app_config.listen_port)

config = uvicorn.Config(
app,
host='0.0.0.0',
Expand All @@ -78,13 +82,19 @@ async def run_udf_app(

# Register the functions only if the app is running interactively.
if app_config.running_interactively:
app.register_functions(replace=True)
app.register_interactive_functions()

_running_server = AwaitableUvicornServer(config)
asyncio.create_task(_running_server.serve())
await _running_server.wait_for_startup()

print(f'Python UDF registered at {base_url}')
if app_config.running_interactively:
sql_names = [
info['signature']['name']
for _func, info in app.endpoints.values()
]
print(f'Registered SQL functions: {", ".join(sql_names)}')

return UdfConnectionInfo(base_url, app.get_function_info())

Expand Down
106 changes: 104 additions & 2 deletions singlestoredb/functions/ext/asgi.py
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,10 @@
from ..signature import signature_to_sql
from ..typing import Masked
from ..typing import Table
from .function_url import classify_interactive_registration
from .function_url import extract_service_url
from .function_url import is_function_not_defined
from .function_url import urls_equal
from .timer import Timer
from singlestoredb.docstring.parser import parse
from singlestoredb.functions.dtypes import escape_name
Expand Down Expand Up @@ -1607,8 +1611,7 @@ def _locate_app_functions(self, cur: Any) -> Tuple[Set[str], Set[str]]:
# See if function URL matches url
cur.execute(f'SHOW CREATE FUNCTION {database_prefix}{escape_name(name)}')
for fname, _, code, *_ in list(cur):
m = re.search(r" (?:\w+) (?:SERVICE|MANAGED) '([^']+)'", code)
if m and m.group(1) == self.url:
if urls_equal(extract_service_url(code), self.url):
funcs.add(f'{database_prefix}{escape_name(fname)}')
if link and re.match(r'^py_ext_func_link_\S{14}$', link):
links.add(link)
Expand Down Expand Up @@ -1812,6 +1815,105 @@ def register_functions(
for func in self.get_create_functions(replace=replace):
cur.execute(func)

def _service_url_for_qualified(
self,
cur: Any,
qualified: str,
) -> Optional[str]:
try:
cur.execute(f'SHOW CREATE FUNCTION {qualified}')
except Exception as exc:
if is_function_not_defined(exc):
return None
raise
rows = list(cur)
if not rows:
return None
code = rows[0][2]
if isinstance(code, bytes):
code = code.decode('utf-8')
return extract_service_url(code)

def _show_create_service_url(self, cur: Any, sql_name: str) -> Optional[str]:
if self.function_database:
qualified = (
f'{escape_name(self.function_database)}.{escape_name(sql_name)}'
)
else:
qualified = escape_name(sql_name)
return self._service_url_for_qualified(cur, qualified)

def _classify_interactive_name(self, cur: Any, sql_name: str) -> str:
existing = self._show_create_service_url(cur, sql_name)
try:
return classify_interactive_registration(existing, self.url)
except ValueError as exc:
raise RuntimeError(
f'Cannot register SQL function `{sql_name}`: {exc}',
) from exc

def preflight_interactive_functions(
self,
*connection_args: Any,
**connection_kwargs: Any,
) -> None:
"""Raise if any current name is published or owned by another session.

Read-only: does not CREATE or DROP functions.
"""
with connection.connect(*connection_args, **connection_kwargs) as conn:
with conn.cursor() as cur:
for _key, (_endpoint, info) in self.endpoints.items():
sql_name = info['signature']['name']
self._classify_interactive_name(cur, sql_name)

def register_interactive_functions(
self,
*connection_args: Any,
**connection_kwargs: Any,
) -> None:
Comment thread
KarishS2 marked this conversation as resolved.
"""Register functions for an interactive notebook session.

Creates or replaces a name only when it is missing or already
points at this session's interactive URL. Published and
other-session functions are left unchanged. Ownership is
re-checked immediately before each write or drop.
"""
with connection.connect(*connection_args, **connection_kwargs) as conn:
with conn.cursor() as cur:
if self.function_database:
database_prefix = escape_name(self.function_database) + '.'
else:
database_prefix = ''
current_names = set()
for _key, (_endpoint, info) in self.endpoints.items():
sql_name = info['signature']['name']
current_names.add(
f'{database_prefix}{escape_name(sql_name)}',
)
self._classify_interactive_name(cur, sql_name)

funcs, _links = self._locate_app_functions(cur)
for fname in funcs:
if fname not in current_names:
existing = self._service_url_for_qualified(cur, fname)
if urls_equal(existing, self.url):
cur.execute(f'DROP FUNCTION IF EXISTS {fname}')

for _key, (_endpoint, info) in self.endpoints.items():
sig = info['signature']
action = self._classify_interactive_name(cur, sig['name'])
cur.execute(
signature_to_sql(
sig,
url=self.url,
data_format=self.data_format,
app_mode=self.app_mode,
replace=(action == 'replace'),
database=self.function_database or None,
),
)

def drop_functions(
self,
*connection_args: Any,
Expand Down
72 changes: 72 additions & 0 deletions singlestoredb/functions/ext/function_url.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,72 @@
"""Parse and classify external function service URLs from SHOW CREATE."""
from __future__ import annotations

import re
from typing import Optional

from ...mysql.constants import ER


# Managed Python UDFs emit: AS MANAGED SERVICE "https://..."
# Older / remote forms may use SERVICE 'https://...'
_SERVICE_URL_RE = re.compile(
r'(?:MANAGED\s+)?SERVICE\s+[\'"]([^\'"]+)[\'"]',
re.IGNORECASE,
)


def extract_service_url(create_sql: str) -> Optional[str]:
"""Return the MANAGED/REMOTE SERVICE URL from SHOW CREATE FUNCTION text."""
if not create_sql:
return None
match = _SERVICE_URL_RE.search(create_sql)
if match is None:
return None
return match.group(1)


def normalize_service_url(url: str) -> str:
return url.rstrip('/')


def urls_equal(left: Optional[str], right: Optional[str]) -> bool:
if not left or not right:
return False
return normalize_service_url(left) == normalize_service_url(right)


def is_interactive_service_url(url: Optional[str]) -> bool:
if not url:
return False
return '/interactive' in normalize_service_url(url).lower()


def classify_interactive_registration(
existing_url: Optional[str],
this_session_url: str,
) -> str:
"""Return 'create', 'replace', or raise ValueError if the name is not ours.

Interactive registration may only create a missing name or replace a
function that already points at this notebook session's /interactive/ URL.
"""
if not existing_url:
return 'create'
if urls_equal(existing_url, this_session_url):
return 'replace'
raise ValueError(
f'Cannot register over existing function pointing at {existing_url} '
f'(this session is {this_session_url}). '
'Interactive registration will not replace a published or '
'other-session function.',
)


def is_function_not_defined(exc: BaseException) -> bool:
errno = getattr(exc, 'errno', None)
if errno in (ER.FUNCTION_NOT_DEFINED, ER.SP_DOES_NOT_EXIST):
return True
args = getattr(exc, 'args', ())
if args and args[0] in (ER.FUNCTION_NOT_DEFINED, ER.SP_DOES_NOT_EXIST):
return True
return False
Comment thread
cursor[bot] marked this conversation as resolved.
Loading
Loading