mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-15 23:31:29 +00:00
fix(rust): type driver initialization errors
This commit is contained in:
parent
0688d9e309
commit
aaace234bb
3 changed files with 54 additions and 6 deletions
|
|
@ -2,6 +2,8 @@ use std::ffi::CString;
|
|||
|
||||
use pyo3::prelude::*;
|
||||
|
||||
use crate::errors::RustBridgeDriverError;
|
||||
|
||||
const DRIVE: &str = r#"
|
||||
def drive_sync(arguments):
|
||||
host = Host(arguments, False)
|
||||
|
|
@ -37,12 +39,46 @@ pub(crate) fn compile<'py>(
|
|||
route: &str,
|
||||
host: &str,
|
||||
) -> PyResult<Bound<'py, PyModule>> {
|
||||
let source = CString::new(format!("{host}\n{DRIVE}")).map_err(|_| {
|
||||
pyo3::exceptions::PyValueError::new_err("driver source contains a null byte")
|
||||
})?;
|
||||
let source = CString::new(format!("{host}\n{DRIVE}"))
|
||||
.map_err(|_| RustBridgeDriverError::new_err("driver source contains a null byte"))?;
|
||||
let filename = CString::new(format!("{route}_driver.py"))
|
||||
.map_err(|_| pyo3::exceptions::PyValueError::new_err("invalid driver route name"))?;
|
||||
.map_err(|_| RustBridgeDriverError::new_err("driver route name contains a null byte"))?;
|
||||
let module_name = CString::new(format!("_{route}_driver"))
|
||||
.map_err(|_| pyo3::exceptions::PyValueError::new_err("invalid driver route name"))?;
|
||||
.map_err(|_| RustBridgeDriverError::new_err("driver route name contains a null byte"))?;
|
||||
PyModule::from_code(py, &source, &filename, &module_name)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn null_byte_in_route_raises_driver_error() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let error = compile(py, "invalid\0route", "class Host: pass")
|
||||
.expect_err("route names containing null bytes should fail");
|
||||
|
||||
assert!(error.is_instance_of::<RustBridgeDriverError>(py));
|
||||
assert_eq!(
|
||||
error.to_string(),
|
||||
"RustBridgeDriverError: driver route name contains a null byte"
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn null_byte_in_source_raises_driver_error() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let error = compile(py, "test", "class Host:\0 pass")
|
||||
.expect_err("driver source containing null bytes should fail");
|
||||
|
||||
assert!(error.is_instance_of::<RustBridgeDriverError>(py));
|
||||
assert_eq!(
|
||||
error.to_string(),
|
||||
"RustBridgeDriverError: driver source contains a null byte"
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -16,6 +16,13 @@ pyo3::create_exception!(
|
|||
"The provider call was already issued and failed. Args are (status, message); status is 0 when there was no HTTP response."
|
||||
);
|
||||
|
||||
pyo3::create_exception!(
|
||||
_native,
|
||||
RustBridgeDriverError,
|
||||
pyo3::exceptions::PyRuntimeError,
|
||||
"The Rust bridge could not initialize a route driver."
|
||||
);
|
||||
|
||||
pub(crate) fn core_error_to_pyerr(err: Error) -> PyErr {
|
||||
match err {
|
||||
Error::InvalidProvider(_) => PyValueError::new_err("Invalid provider configuration"),
|
||||
|
|
@ -81,7 +88,11 @@ pub(crate) fn messages_provider_error_to_pyerr(err: Error) -> PyErr {
|
|||
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
let py = module.py();
|
||||
module.add("RustBridgeDeclined", py.get_type::<RustBridgeDeclined>())?;
|
||||
module.add("RustUpstreamError", py.get_type::<RustUpstreamError>())
|
||||
module.add("RustUpstreamError", py.get_type::<RustUpstreamError>())?;
|
||||
module.add(
|
||||
"RustBridgeDriverError",
|
||||
py.get_type::<RustBridgeDriverError>(),
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
|
|
|||
|
|
@ -35,6 +35,7 @@ mod tests {
|
|||
let expected = [
|
||||
"RustBridgeDeclined",
|
||||
"RustUpstreamError",
|
||||
"RustBridgeDriverError",
|
||||
"ocr",
|
||||
"aocr",
|
||||
"transcription",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue