|
32 | 32 | from pydantic import ValidationError |
33 | 33 | from pydantic_ai import capture_run_messages |
34 | 34 | from pydantic_ai.messages import ModelMessagesTypeAdapter |
35 | | -from sqlalchemy.engine import make_url |
| 35 | +from pyobvector import AsyncOceanBaseDialect |
| 36 | +from sqlalchemy.engine import URL, make_url |
| 37 | +from sqlalchemy.exc import ArgumentError |
36 | 38 |
|
37 | 39 | from evaluation.memory.locomo.dataset import LoCoMoSession |
38 | 40 | from evaluation.memory.locomo.metrics import retrieval_metrics |
@@ -280,28 +282,68 @@ def _database_config(settings: ServerSettings, output_directory: Path) -> Databa |
280 | 282 | return settings.database |
281 | 283 |
|
282 | 284 |
|
| 285 | +def _oceanbase_target(url: URL) -> dict[str, Any]: |
| 286 | + if any(not isinstance(value, str) for value in url.query.values()): |
| 287 | + raise ValueError("OceanBase URL query parameters must not be repeated") # noqa: TRY003 |
| 288 | + if {key.casefold() for key in url.query} & {"read_default_file", "read_default_group", "sql_mode"}: |
| 289 | + raise ValueError("OceanBase external defaults and SQL mode overrides are not supported for resumable runs") # noqa: TRY003 |
| 290 | + try: |
| 291 | + # URL query values override the translated authority/path in the installed dialect. |
| 292 | + _, options = AsyncOceanBaseDialect().create_connect_args(url) |
| 293 | + username = options.get("user") |
| 294 | + database = options.get("db") |
| 295 | + host = options.get("host") |
| 296 | + port = int(options.get("port", 3306)) |
| 297 | + if ( |
| 298 | + not isinstance(host, str) |
| 299 | + or not host |
| 300 | + or not all(isinstance(value, str) and value for value in (username, database)) |
| 301 | + or not 1 <= port <= 65535 |
| 302 | + ): |
| 303 | + raise ValueError # noqa: TRY301 |
| 304 | + except (ArgumentError, TypeError, ValueError): |
| 305 | + raise ValueError( # noqa: TRY003 |
| 306 | + "OceanBase connection target must explicitly identify its host, port, user and database" |
| 307 | + ) from None |
| 308 | + socket = options.get("unix_socket") |
| 309 | + # Passwords, TLS/authentication material and tuning options do not identify the database. |
| 310 | + # init_command is replaced by the runtime's fixed SET autocommit = 0 command. |
| 311 | + return { |
| 312 | + "driver": url.drivername, |
| 313 | + "username": username, |
| 314 | + "database": database, |
| 315 | + "host": None if socket else host.lower(), |
| 316 | + "port": None if socket else port, |
| 317 | + "unix_socket": str(Path(socket).resolve()) if socket else None, |
| 318 | + } |
| 319 | + |
| 320 | + |
283 | 321 | def _database_identity(database: DatabaseConfig) -> dict[str, str]: |
| 322 | + if database.kind == "oceanbase": |
| 323 | + target = _oceanbase_target(make_url(database.url.get_secret_value())) |
| 324 | + return { |
| 325 | + "database_kind": database.kind, |
| 326 | + "database_fingerprint_version": "oceanbase-target-v2", |
| 327 | + "database_fingerprint": _digest(json.dumps(target, sort_keys=True)), |
| 328 | + } |
284 | 329 | target: dict[str, Any] |
285 | 330 | if database.kind == "seekdb": |
286 | 331 | target = {"path": str(database.path.resolve()), "database": database.database} |
287 | 332 | else: |
288 | | - value = database.url if database.kind == "sqlite" else database.url.get_secret_value() |
289 | | - url = make_url(value) |
290 | | - if database.kind == "sqlite": |
291 | | - name = url.database |
292 | | - if name and name != ":memory:" and not name.startswith("file:"): |
293 | | - name = str(Path(name).resolve()) |
294 | | - url = url.set(database=name) |
295 | | - # The OceanBase username can select a tenant. Password rotation does not change database identity. |
| 333 | + url = make_url(database.url) |
| 334 | + name = url.database |
| 335 | + if name and name != ":memory:" and not name.startswith("file:"): |
| 336 | + name = str(Path(name).resolve()) |
| 337 | + url = url.set(database=name) |
296 | 338 | target = { |
297 | 339 | "driver": url.drivername, |
298 | 340 | "host": (url.host or "").lower(), |
299 | | - "port": url.port or (3306 if database.kind == "oceanbase" else None), |
| 341 | + "port": url.port or None, |
300 | 342 | "username": url.username, |
301 | 343 | "database": url.database, |
302 | 344 | "query": dict(url.query), |
303 | 345 | } |
304 | | - if database.kind == "sqlite" and (url.database or "").startswith("file:"): |
| 346 | + if (url.database or "").startswith("file:"): |
305 | 347 | target["working_directory"] = str(Path.cwd().resolve()) |
306 | 348 | return { |
307 | 349 | "database_kind": database.kind, |
|
0 commit comments