From bdb997c81d7aedee533c5aaf2fd3a04644b04a5e Mon Sep 17 00:00:00 2001 From: Christopher Rowley Date: Sat, 3 Oct 2026 14:37:02 +0100 Subject: [PATCH] Add zero-copy PyString wrapper --- src/API/exports.jl | 1 + src/API/types.jl | 15 ++++++ src/C/pointers.jl | 1 + src/Utils/Utils.jl | 114 ++++++++++++++++++++----------------------- src/Wrap/PyString.jl | 21 ++++++++ src/Wrap/Wrap.jl | 4 +- test/Convert.jl | 22 +++++++++ test/Utils.jl | 9 ++++ 8 files changed, 125 insertions(+), 62 deletions(-) create mode 100644 src/Wrap/PyString.jl diff --git a/src/API/exports.jl b/src/API/exports.jl index 3d476182..5fec736d 100644 --- a/src/API/exports.jl +++ b/src/API/exports.jl @@ -119,6 +119,7 @@ export PyIterable export PyList export PyPandasDataFrame export PySet +export PyString export PyTable # JlWrap diff --git a/src/API/types.jl b/src/API/types.jl index 55a4211e..99cb5c24 100644 --- a/src/API/types.jl +++ b/src/API/types.jl @@ -102,6 +102,21 @@ struct PyDict{K,V} <: AbstractDict{K,V} PyDict{K,V}(x = pydict()) where {K,V} = new{K,V}(ispy(x) ? Py(x) : pydict(x)) end +""" + PyString(x) + +Wrap the Python `str` object `x` as an `AbstractString` without copying its UTF-8 data. + +If `x` is not a Python object, it is first converted to a Python `str`. +""" +struct PyString <: AbstractString + py::Py + ptr::Ptr{UInt8} + length::Int + PyString(::Val{:new}, py::Py, ptr::Ptr{UInt8}, length::Int) = + new(py, ptr, length) +end + """ PyIO(x; own=false, text=missing, line_buffering=false, buflen=4096) diff --git a/src/C/pointers.jl b/src/C/pointers.jl index 9644329f..7171f36a 100644 --- a/src/C/pointers.jl +++ b/src/C/pointers.jl @@ -141,6 +141,7 @@ const CAPI_FUNC_SIGS = Dict{Symbol,Pair{Tuple,Type}}( :PyComplex_AsCComplex => (PyPtr,) => Py_complex, # STR :PyUnicode_DecodeUTF8 => (Ptr{Cchar}, Py_ssize_t, Ptr{Cchar}) => PyPtr, + :PyUnicode_AsUTF8AndSize => (PyPtr, Ptr{Py_ssize_t}) => Ptr{Cchar}, :PyUnicode_AsUTF8String => (PyPtr,) => PyPtr, :PyUnicode_InternInPlace => (Ptr{PyPtr},) => Cvoid, # BYTES diff --git a/src/Utils/Utils.jl b/src/Utils/Utils.jl index a26b3198..2c600b49 100644 --- a/src/Utils/Utils.jl +++ b/src/Utils/Utils.jl @@ -231,82 +231,74 @@ Base.codeunit(x::StaticString, i::Integer) = x.codeunits[i] Base.codeunit(x::StaticString{T}) where {T} = T -function Base.isvalid(x::StaticString{UInt8,N}, i::Int) where {N} - if i < 1 || i > N - return false - end - cs = x.codeunits - c = @inbounds cs[i] - if all(iszero, (cs[j] for j = i:N)) - return false - elseif (c & 0x80) == 0x00 +function utf8_isvalid(x, n::Int, i::Int) + 1 ≤ i ≤ n || return false + c = @inbounds codeunit(x, i) + if (c & 0x80) == 0x00 return true elseif (c & 0x40) == 0x00 return false elseif (c & 0x20) == 0x00 - return @inbounds (i ≤ N - 1) && ((cs[i+1] & 0xC0) == 0x80) + return @inbounds (i ≤ n - 1) && ((codeunit(x, i + 1) & 0xC0) == 0x80) elseif (c & 0x10) == 0x00 - return @inbounds (i ≤ N - 2) && - ((cs[i+1] & 0xC0) == 0x80) && - ((cs[i+2] & 0xC0) == 0x80) + return @inbounds (i ≤ n - 2) && + ((codeunit(x, i + 1) & 0xC0) == 0x80) && + ((codeunit(x, i + 2) & 0xC0) == 0x80) elseif (c & 0x08) == 0x00 - return @inbounds (i ≤ N - 3) && - ((cs[i+1] & 0xC0) == 0x80) && - ((cs[i+2] & 0xC0) == 0x80) && - ((cs[i+3] & 0xC0) == 0x80) + return @inbounds (i ≤ n - 3) && + ((codeunit(x, i + 1) & 0xC0) == 0x80) && + ((codeunit(x, i + 2) & 0xC0) == 0x80) && + ((codeunit(x, i + 3) & 0xC0) == 0x80) else return false end - return false end -function Base.iterate(x::StaticString{UInt8,N}, i::Int = 1) where {N} - i > N && return - cs = x.codeunits - c = @inbounds cs[i] - if all(iszero, (cs[j] for j = i:N)) - return - elseif (c & 0x80) == 0x00 +function utf8_iterate(x, n::Int, i::Int = 1) + i > n && return nothing + utf8_isvalid(x, n, i) || throw(StringIndexError(x, i)) + c = @inbounds codeunit(x, i) + if (c & 0x80) == 0x00 return (reinterpret(Char, UInt32(c) << 24), i + 1) - elseif (c & 0x40) == 0x00 - nothing elseif (c & 0x20) == 0x00 - if @inbounds (i ≤ N - 1) && ((cs[i+1] & 0xC0) == 0x80) - return ( - reinterpret(Char, (UInt32(cs[i]) << 24) | (UInt32(cs[i+1]) << 16)), - i + 2, - ) - end + return ( + reinterpret(Char, (UInt32(c) << 24) | (UInt32(codeunit(x, i + 1)) << 16)), + i + 2, + ) elseif (c & 0x10) == 0x00 - if @inbounds (i ≤ N - 2) && ((cs[i+1] & 0xC0) == 0x80) && ((cs[i+2] & 0xC0) == 0x80) - return ( - reinterpret( - Char, - (UInt32(cs[i]) << 24) | - (UInt32(cs[i+1]) << 16) | - (UInt32(cs[i+2]) << 8), - ), - i + 3, - ) - end - elseif (c & 0x08) == 0x00 - if @inbounds (i ≤ N - 3) && - ((cs[i+1] & 0xC0) == 0x80) && - ((cs[i+2] & 0xC0) == 0x80) && - ((cs[i+3] & 0xC0) == 0x80) - return ( - reinterpret( - Char, - (UInt32(cs[i]) << 24) | - (UInt32(cs[i+1]) << 16) | - (UInt32(cs[i+2]) << 8) | - UInt32(cs[i+3]), - ), - i + 4, - ) - end + return ( + reinterpret( + Char, + (UInt32(c) << 24) | + (UInt32(codeunit(x, i + 1)) << 16) | + (UInt32(codeunit(x, i + 2)) << 8), + ), + i + 3, + ) + else + return ( + reinterpret( + Char, + (UInt32(c) << 24) | + (UInt32(codeunit(x, i + 1)) << 16) | + (UInt32(codeunit(x, i + 2)) << 8) | + UInt32(codeunit(x, i + 3)), + ), + i + 4, + ) end - throw(StringIndexError(x, i)) +end + +function Base.isvalid(x::StaticString{UInt8,N}, i::Int) where {N} + cs = x.codeunits + n = something(findlast(!iszero, cs), 0) + return utf8_isvalid(x, n, i) +end + +function Base.iterate(x::StaticString{UInt8,N}, i::Int = 1) where {N} + cs = x.codeunits + n = something(findlast(!iszero, cs), 0) + return utf8_iterate(x, n, i) end function Base.isvalid(x::StaticString{UInt32,N}, i::Int) where {N} diff --git a/src/Wrap/PyString.jl b/src/Wrap/PyString.jl new file mode 100644 index 00000000..10af9001 --- /dev/null +++ b/src/Wrap/PyString.jl @@ -0,0 +1,21 @@ +function PyString(x) + py = ispy(x) ? Py(x) : pystr(x) + n = Ref{C.Py_ssize_t}() + ptr = C.PyUnicode_AsUTF8AndSize(py, n) + ptr == C_NULL && pythrow() + return PyString(Val(:new), py, Ptr{UInt8}(ptr), Int(n[])) +end + +ispy(::PyString) = true +Py(x::PyString) = x.py + +pyconvert_rule_string(::Type{PyString}, x::Py) = pyconvert_return(PyString(x)) + +Base.ncodeunits(x::PyString) = x.length +Base.codeunit(::PyString) = UInt8 +Base.@propagate_inbounds function Base.codeunit(x::PyString, i::Integer) + @boundscheck checkbounds(1:x.length, i) + return unsafe_load(x.ptr, i) +end +Base.isvalid(x::PyString, i::Int) = Utils.utf8_isvalid(x, x.length, i) +Base.iterate(x::PyString, i::Int = 1) = Utils.utf8_iterate(x, x.length, i) diff --git a/src/Wrap/Wrap.jl b/src/Wrap/Wrap.jl index d3b30ff2..e3ceed65 100644 --- a/src/Wrap/Wrap.jl +++ b/src/Wrap/Wrap.jl @@ -14,7 +14,7 @@ using ..Convert using ..PyMacro import ..PythonCall: - PyArray, PyDict, PyIO, PyIterable, PyList, PyPandasDataFrame, PySet, PyTable + PyArray, PyDict, PyIO, PyIterable, PyList, PyPandasDataFrame, PySet, PyString, PyTable using Base: @propagate_inbounds using Tables: Tables @@ -26,6 +26,7 @@ include("PyIterable.jl") include("PyDict.jl") include("PyList.jl") include("PySet.jl") +include("PyString.jl") include("PyArray.jl") include("PyIO.jl") include("PyTable.jl") @@ -73,6 +74,7 @@ function __init__() end priority = PYCONVERT_PRIORITY_NORMAL + pyconvert_add_rule("builtins:str", PyString, pyconvert_rule_string, priority) pyconvert_add_rule("", Array, pyconvert_rule_array, priority) pyconvert_add_rule("", Array, pyconvert_rule_array, priority) pyconvert_add_rule("", Array, pyconvert_rule_array, priority) diff --git a/test/Convert.jl b/test/Convert.jl index 2765440f..6575202a 100644 --- a/test/Convert.jl +++ b/test/Convert.jl @@ -93,6 +93,28 @@ end @test x2 === "αβγℵ√" end +@testitem "str → PyString" begin + x = pystr("aβℵ\0🙂") + s = pyconvert(PyString, x) + @test s isa PyString + @test Py(s) === s.py + @test String(s) == "aβℵ\0🙂" + @test ncodeunits(s) == 11 + @test codeunit(s) === UInt8 + @test collect(s) == ['a', 'β', 'ℵ', '\0', '🙂'] + @test isvalid(s, 1) + @test !isvalid(s, 3) + @test codeunit(s, 3) == 0xb2 + @test_throws StringIndexError s[3] + @test_throws BoundsError codeunit(s, 12) + + t = PyString("hello") + @test String(t) == "hello" + @test isempty(PyString("")) + @test_throws PyException PyString(pyint(1)) + @test pyconvert(Any, x) isa String +end + @testitem "str → Symbol" begin x1 = pyconvert(Symbol, pystr("hello")) @test x1 === :hello diff --git a/test/Utils.jl b/test/Utils.jl index 33c82bd9..be3e95fe 100644 --- a/test/Utils.jl +++ b/test/Utils.jl @@ -22,3 +22,12 @@ end @test s[1:2] == "ab" @test s[1:2:end] == "aaaab" end + +@testitem "StaticString UTF-8" begin + S = PythonCall.Utils.StaticString + s = S{UInt8,10}("aβℵ🙂") + @test String(s) == "aβℵ🙂" + @test collect(s) == ['a', 'β', 'ℵ', '🙂'] + @test !isvalid(S{UInt8,1}((0x80,)), 1) + @test !isvalid(S{UInt8,1}((0xf8,)), 1) +end