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
7 changes: 5 additions & 2 deletions .github/workflows/pr-test-rust.yml
Original file line number Diff line number Diff line change
Expand Up @@ -419,6 +419,7 @@ jobs:
run: |
bash scripts/ci_agentic_svc_deps.sh setup-oracle-client
bash scripts/ci_agentic_svc_deps.sh create-oracle-user oracle-db
bash scripts/ci_agentic_svc_deps.sh create-oracle-flyway-user oracle-db

- name: Run E2E tests
env:
Expand All @@ -431,9 +432,11 @@ jobs:
bash scripts/ci_killall_sglang.sh "nuk_gpus"
${{ matrix.env_vars }} ROUTER_LOCAL_MODEL_PATH="/home/ubuntu/models" pytest ${{ matrix.reruns }} ${{ matrix.parallel_opts }} ${{ matrix.ignore_opts }} ${{ matrix.test_dirs }} ${{ matrix.test_filter }} -s -vv -o log_cli=true --log-cli-level=INFO

- name: Cleanup Oracle test user
- name: Cleanup Oracle test users
if: always() && matrix.setup_agentic_deps
run: bash scripts/ci_agentic_svc_deps.sh cleanup-oracle-user oracle-db
run: |
bash scripts/ci_agentic_svc_deps.sh cleanup-oracle-flyway-user oracle-db
bash scripts/ci_agentic_svc_deps.sh cleanup-oracle-user oracle-db

- name: Upload benchmark results
if: matrix.upload_benchmarks && success()
Expand Down
1 change: 1 addition & 0 deletions bindings/python/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ crate-type = ["cdylib"]
pyo3 = { version = "0.28.2", features = ["extension-module", "abi3-py38"] }
tokio = { version = "1.42.0", features = ["full"] }
once_cell = "1.19"
serde_yaml = "0.9"

[dependencies.smg]
path = "../../model_gateway"
Expand Down
43 changes: 36 additions & 7 deletions bindings/python/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -452,6 +452,7 @@ struct Router {
enable_trace: bool,
otlp_traces_endpoint: String,
control_plane_auth: Option<PyControlPlaneAuthConfig>,
schema_config: Option<String>,
}

impl Router {
Expand Down Expand Up @@ -587,24 +588,49 @@ impl Router {
HistoryBackendType::Redis => config::HistoryBackend::Redis,
};

// Load schema config from YAML file if provided
let schema = if let Some(ref path) = self.schema_config {
let content = std::fs::read_to_string(path).map_err(|e| {
config::ConfigError::ValidationFailed {
reason: format!("Failed to read schema config file '{path}': {e}"),
}
})?;
let schema: config::SchemaConfig = serde_yaml::from_str(&content).map_err(|e| {
config::ConfigError::ValidationFailed {
reason: format!("Failed to parse schema config file '{path}': {e}"),
}
})?;
Some(schema)
} else {
None
};
Comment on lines +592 to +606

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

This schema loading logic is duplicated in model_gateway/src/main.rs (inside load_schema_config). To improve maintainability and avoid code duplication, consider extracting this logic into a shared function. A good place for this would be within the smg::config module, for example, as pub fn load_schema_from_path(path: &str) -> ConfigResult<SchemaConfig>.

References
  1. Extract duplicated logic into a shared helper function to improve maintainability and reduce redundancy.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Appreciate the suggestion, but these two callers live in separate entry-point crates (CLI binary vs Python bindings). The logic is ~6 lines of straightforward read-file-then-parse glue — extracting it would add public API surface to the config module for something only used at initialization. Keeping it inline in each caller is fine here.


let oracle = if matches!(self.history_backend, HistoryBackendType::Oracle) {
self.oracle_config
.as_ref()
.map(|cfg| cfg.to_config_oracle())
self.oracle_config.as_ref().map(|cfg| {
let mut c = cfg.to_config_oracle();
c.schema.clone_from(&schema);
c
})
} else {
None
};

let postgres_config = if matches!(self.history_backend, HistoryBackendType::Postgres) {
self.postgres_config
.as_ref()
.map(|cfg| cfg.to_config_postgres())
self.postgres_config.as_ref().map(|cfg| {
let mut c = cfg.to_config_postgres();
c.schema.clone_from(&schema);
c
})
} else {
None
};

let redis_config = if matches!(self.history_backend, HistoryBackendType::Redis) {
self.redis_config.as_ref().map(|cfg| cfg.to_config_redis())
self.redis_config.as_ref().map(|cfg| {
let mut c = cfg.to_config_redis();
c.schema = schema;
c
})
} else {
None
};
Expand Down Expand Up @@ -779,6 +805,7 @@ impl Router {
enable_trace = false,
otlp_traces_endpoint = String::from("localhost:4317"),
control_plane_auth = None,
schema_config = None,
))]
#[expect(clippy::too_many_arguments)]
#[expect(
Expand Down Expand Up @@ -874,6 +901,7 @@ impl Router {
enable_trace: bool,
otlp_traces_endpoint: String,
control_plane_auth: Option<PyControlPlaneAuthConfig>,
schema_config: Option<String>,
) -> PyResult<Self> {
let mut all_urls = worker_urls.clone();

Expand Down Expand Up @@ -979,6 +1007,7 @@ impl Router {
enable_trace,
otlp_traces_endpoint,
control_plane_auth,
schema_config,
})
}

Expand Down
9 changes: 9 additions & 0 deletions bindings/python/src/smg/router_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -129,6 +129,7 @@ class RouterArgs:
redis_url: str | None = None
redis_pool_max: int = 16
redis_retention_days: int = 30
schema_config: str | None = None
# mTLS configuration for worker communication
client_cert_path: str | None = None
client_key_path: str | None = None
Expand Down Expand Up @@ -871,6 +872,14 @@ def add_cli_args(
help="Redis data retention in days (-1 for persistent, default: 30, env: REDIS_RETENTION_DAYS)",
)

# Schema configuration
backend_group.add_argument(
f"--{prefix}schema-config",
type=str,
default=None,
help="Path to a YAML schema config file for storage table/column remapping",
)

# TLS/mTLS configuration
tls_group.add_argument(
f"--{prefix}client-cert-path",
Expand Down
2 changes: 1 addition & 1 deletion e2e_test/fixtures/hooks.py
Original file line number Diff line number Diff line change
Expand Up @@ -463,7 +463,7 @@ def pytest_configure(config: pytest.Config) -> None:
config.addinivalue_line(
"markers",
"storage(backend): mark test to use a specific history storage backend "
"(memory, oracle). Default is memory.",
"(memory, oracle, oracle-custom). Default is memory.",
)
config.addinivalue_line(
"markers",
Expand Down
32 changes: 29 additions & 3 deletions e2e_test/infra/gateway.py
Original file line number Diff line number Diff line change
Expand Up @@ -155,7 +155,7 @@ def start(
decode_workers: List of decode ModelInstance objects for PD mode.
igw_mode: Start in IGW mode (no workers, add via API).
cloud_backend: Cloud backend type ("openai", "xai", or "anthropic").
history_backend: History backend for cloud mode ("memory" or "oracle").
history_backend: History backend for cloud mode ("memory", "oracle", or "oracle-custom").
policy: Routing policy (round_robin, random, etc.)
timeout: Startup timeout in seconds.
show_output: Show subprocess output (env var override).
Expand Down Expand Up @@ -277,15 +277,41 @@ def start(
raise ValueError(f"Unsupported cloud backend: {cloud_backend}")

backend_type = cloud_backend_type.get(cloud_backend, "openai")

# oracle-custom: use Flyway-managed schema with schema-config
actual_history_backend = history_backend
if history_backend == "oracle-custom":
actual_history_backend = "oracle"
# Override Oracle env vars to use Flyway user credentials
flyway_user = os.environ.get("ATP_FLYWAY_USER", "")
flyway_password = os.environ.get("ATP_FLYWAY_PASSWORD", "")
flyway_dsn = os.environ.get("ATP_FLYWAY_DSN", "")
if not all([flyway_user, flyway_password, flyway_dsn]):
raise ValueError(
"ATP_FLYWAY_USER, ATP_FLYWAY_PASSWORD, and ATP_FLYWAY_DSN "
"environment variables required for oracle-custom backend"
)
self._env["ATP_USER"] = flyway_user
self._env["ATP_PASSWORD"] = flyway_password
self._env["ATP_DSN"] = flyway_dsn

mode_args = [
"--backend",
backend_type,
"--worker-urls",
worker_url,
"--history-backend",
history_backend,
actual_history_backend,
]

if history_backend == "oracle-custom":
mode_args.extend(
[
"--schema-config",
"scripts/oracle_flyway/schema-config.yaml",
]
)
Comment thread
coderabbitai[bot] marked this conversation as resolved.

self._launch(
mode_args=mode_args,
timeout=timeout,
Expand Down Expand Up @@ -621,7 +647,7 @@ def launch_cloud_gateway(

Args:
runtime: Cloud runtime ("openai" or "xai")
history_backend: History storage backend ("memory" or "oracle")
history_backend: History storage backend ("memory", "oracle", or "oracle-custom")
extra_args: Additional router arguments
timeout: Startup timeout in seconds
show_output: Show subprocess output
Expand Down
42 changes: 38 additions & 4 deletions e2e_test/responses/test_state_management.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,13 +18,12 @@


# =============================================================================
# Cloud Backend Tests (OpenAI, xAI)
# Cloud Backend Tests (base mixin — no parametrize, subclasses add their own)
# =============================================================================


@pytest.mark.parametrize("setup_backend", ["openai", "xai"], indirect=True)
class TestStateManagementCloud:
"""State management tests against cloud APIs."""
class _StateManagementCloudBase:
"""Base test methods for state management against cloud APIs."""

def test_basic_response_creation(self, setup_backend, smg):
"""Test basic response creation without state."""
Expand Down Expand Up @@ -185,6 +184,41 @@ def test_mutually_exclusive_parameters(self, setup_backend, smg):
)


# =============================================================================
# Cloud Backend Tests (OpenAI)
# =============================================================================


@pytest.mark.parametrize("setup_backend", ["openai"], indirect=True)
class TestStateManagementCloud(_StateManagementCloudBase):
"""State management tests against OpenAI cloud API."""


# =============================================================================
# Cloud Backend Tests (xAI)
# =============================================================================


@pytest.mark.parametrize("setup_backend", ["xai"], indirect=True)
class TestStateManagementCloudXai(_StateManagementCloudBase):
"""State management tests against xAI cloud API."""


# =============================================================================
# Cloud Backend Tests with Flyway-managed Oracle schema (oracle-custom)
# =============================================================================


@pytest.mark.storage("oracle-custom")
@pytest.mark.parametrize("setup_backend", ["openai"], indirect=True)
class TestStateManagementOracleCustom(_StateManagementCloudBase):
"""State management tests against Oracle with Flyway-managed schema (schema-config).

The storage("oracle-custom") marker causes the gateway to launch with
--schema-config pointing to the Flyway schema, using ATP_FLYWAY_* env vars.
"""


# =============================================================================
# Local Backend Tests (gRPC with Qwen model)
# =============================================================================
Expand Down
4 changes: 3 additions & 1 deletion model_gateway/src/config/types.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,9 @@ use std::collections::HashMap;
use openai_protocol::worker::HealthCheckConfig as ProtocolHealthCheckConfig;
use serde::{Deserialize, Serialize};
// Re-export storage config types from data_connector
pub use smg_data_connector::{HistoryBackend, OracleConfig, PostgresConfig, RedisConfig};
pub use smg_data_connector::{
HistoryBackend, OracleConfig, PostgresConfig, RedisConfig, SchemaConfig,
};

use super::{validation::ConfigValidator, ConfigResult};
use crate::core::ConnectionMode;
Expand Down
Loading
Loading