Skip to content

Commit bdb997c

Browse files
committed
Add zero-copy PyString wrapper
1 parent de68e38 commit bdb997c

8 files changed

Lines changed: 125 additions & 62 deletions

File tree

‎src/API/exports.jl‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -119,6 +119,7 @@ export PyIterable
119119
export PyList
120120
export PyPandasDataFrame
121121
export PySet
122+
export PyString
122123
export PyTable
123124

124125
# JlWrap

‎src/API/types.jl‎

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -102,6 +102,21 @@ struct PyDict{K,V} <: AbstractDict{K,V}
102102
PyDict{K,V}(x = pydict()) where {K,V} = new{K,V}(ispy(x) ? Py(x) : pydict(x))
103103
end
104104

105+
"""
106+
PyString(x)
107+
108+
Wrap the Python `str` object `x` as an `AbstractString` without copying its UTF-8 data.
109+
110+
If `x` is not a Python object, it is first converted to a Python `str`.
111+
"""
112+
struct PyString <: AbstractString
113+
py::Py
114+
ptr::Ptr{UInt8}
115+
length::Int
116+
PyString(::Val{:new}, py::Py, ptr::Ptr{UInt8}, length::Int) =
117+
new(py, ptr, length)
118+
end
119+
105120
"""
106121
PyIO(x; own=false, text=missing, line_buffering=false, buflen=4096)
107122

‎src/C/pointers.jl‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -141,6 +141,7 @@ const CAPI_FUNC_SIGS = Dict{Symbol,Pair{Tuple,Type}}(
141141
:PyComplex_AsCComplex => (PyPtr,) => Py_complex,
142142
# STR
143143
:PyUnicode_DecodeUTF8 => (Ptr{Cchar}, Py_ssize_t, Ptr{Cchar}) => PyPtr,
144+
:PyUnicode_AsUTF8AndSize => (PyPtr, Ptr{Py_ssize_t}) => Ptr{Cchar},
144145
:PyUnicode_AsUTF8String => (PyPtr,) => PyPtr,
145146
:PyUnicode_InternInPlace => (Ptr{PyPtr},) => Cvoid,
146147
# BYTES

‎src/Utils/Utils.jl‎

Lines changed: 53 additions & 61 deletions
Original file line numberDiff line numberDiff line change
@@ -231,82 +231,74 @@ Base.codeunit(x::StaticString, i::Integer) = x.codeunits[i]
231231

232232
Base.codeunit(x::StaticString{T}) where {T} = T
233233

234-
function Base.isvalid(x::StaticString{UInt8,N}, i::Int) where {N}
235-
if i < 1 || i > N
236-
return false
237-
end
238-
cs = x.codeunits
239-
c = @inbounds cs[i]
240-
if all(iszero, (cs[j] for j = i:N))
241-
return false
242-
elseif (c & 0x80) == 0x00
234+
function utf8_isvalid(x, n::Int, i::Int)
235+
1 ≤ i ≤ n || return false
236+
c = @inbounds codeunit(x, i)
237+
if (c & 0x80) == 0x00
243238
return true
244239
elseif (c & 0x40) == 0x00
245240
return false
246241
elseif (c & 0x20) == 0x00
247-
return @inbounds (i ≤ N - 1) && ((cs[i+1] & 0xC0) == 0x80)
242+
return @inbounds (i ≤ n - 1) && ((codeunit(x, i + 1) & 0xC0) == 0x80)
248243
elseif (c & 0x10) == 0x00
249-
return @inbounds (i ≤ N - 2) &&
250-
((cs[i+1] & 0xC0) == 0x80) &&
251-
((cs[i+2] & 0xC0) == 0x80)
244+
return @inbounds (i ≤ n - 2) &&
245+
((codeunit(x, i + 1) & 0xC0) == 0x80) &&
246+
((codeunit(x, i + 2) & 0xC0) == 0x80)
252247
elseif (c & 0x08) == 0x00
253-
return @inbounds (i ≤ N - 3) &&
254-
((cs[i+1] & 0xC0) == 0x80) &&
255-
((cs[i+2] & 0xC0) == 0x80) &&
256-
((cs[i+3] & 0xC0) == 0x80)
248+
return @inbounds (i ≤ n - 3) &&
249+
((codeunit(x, i + 1) & 0xC0) == 0x80) &&
250+
((codeunit(x, i + 2) & 0xC0) == 0x80) &&
251+
((codeunit(x, i + 3) & 0xC0) == 0x80)
257252
else
258253
return false
259254
end
260-
return false
261255
end
262256

263-
function Base.iterate(x::StaticString{UInt8,N}, i::Int = 1) where {N}
264-
i > N && return
265-
cs = x.codeunits
266-
c = @inbounds cs[i]
267-
if all(iszero, (cs[j] for j = i:N))
268-
return
269-
elseif (c & 0x80) == 0x00
257+
function utf8_iterate(x, n::Int, i::Int = 1)
258+
i > n && return nothing
259+
utf8_isvalid(x, n, i) || throw(StringIndexError(x, i))
260+
c = @inbounds codeunit(x, i)
261+
if (c & 0x80) == 0x00
270262
return (reinterpret(Char, UInt32(c) << 24), i + 1)
271-
elseif (c & 0x40) == 0x00
272-
nothing
273263
elseif (c & 0x20) == 0x00
274-
if @inbounds (i ≤ N - 1) && ((cs[i+1] & 0xC0) == 0x80)
275-
return (
276-
reinterpret(Char, (UInt32(cs[i]) << 24) | (UInt32(cs[i+1]) << 16)),
277-
i + 2,
278-
)
279-
end
264+
return (
265+
reinterpret(Char, (UInt32(c) << 24) | (UInt32(codeunit(x, i + 1)) << 16)),
266+
i + 2,
267+
)
280268
elseif (c & 0x10) == 0x00
281-
if @inbounds (i ≤ N - 2) && ((cs[i+1] & 0xC0) == 0x80) && ((cs[i+2] & 0xC0) == 0x80)
282-
return (
283-
reinterpret(
284-
Char,
285-
(UInt32(cs[i]) << 24) |
286-
(UInt32(cs[i+1]) << 16) |
287-
(UInt32(cs[i+2]) << 8),
288-
),
289-
i + 3,
290-
)
291-
end
292-
elseif (c & 0x08) == 0x00
293-
if @inbounds (i ≤ N - 3) &&
294-
((cs[i+1] & 0xC0) == 0x80) &&
295-
((cs[i+2] & 0xC0) == 0x80) &&
296-
((cs[i+3] & 0xC0) == 0x80)
297-
return (
298-
reinterpret(
299-
Char,
300-
(UInt32(cs[i]) << 24) |
301-
(UInt32(cs[i+1]) << 16) |
302-
(UInt32(cs[i+2]) << 8) |
303-
UInt32(cs[i+3]),
304-
),
305-
i + 4,
306-
)
307-
end
269+
return (
270+
reinterpret(
271+
Char,
272+
(UInt32(c) << 24) |
273+
(UInt32(codeunit(x, i + 1)) << 16) |
274+
(UInt32(codeunit(x, i + 2)) << 8),
275+
),
276+
i + 3,
277+
)
278+
else
279+
return (
280+
reinterpret(
281+
Char,
282+
(UInt32(c) << 24) |
283+
(UInt32(codeunit(x, i + 1)) << 16) |
284+
(UInt32(codeunit(x, i + 2)) << 8) |
285+
UInt32(codeunit(x, i + 3)),
286+
),
287+
i + 4,
288+
)
308289
end
309-
throw(StringIndexError(x, i))
290+
end
291+
292+
function Base.isvalid(x::StaticString{UInt8,N}, i::Int) where {N}
293+
cs = x.codeunits
294+
n = something(findlast(!iszero, cs), 0)
295+
return utf8_isvalid(x, n, i)
296+
end
297+
298+
function Base.iterate(x::StaticString{UInt8,N}, i::Int = 1) where {N}
299+
cs = x.codeunits
300+
n = something(findlast(!iszero, cs), 0)
301+
return utf8_iterate(x, n, i)
310302
end
311303

312304
function Base.isvalid(x::StaticString{UInt32,N}, i::Int) where {N}

‎src/Wrap/PyString.jl‎

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,21 @@
1+
function PyString(x)
2+
py = ispy(x) ? Py(x) : pystr(x)
3+
n = Ref{C.Py_ssize_t}()
4+
ptr = C.PyUnicode_AsUTF8AndSize(py, n)
5+
ptr == C_NULL && pythrow()
6+
return PyString(Val(:new), py, Ptr{UInt8}(ptr), Int(n[]))
7+
end
8+
9+
ispy(::PyString) = true
10+
Py(x::PyString) = x.py
11+
12+
pyconvert_rule_string(::Type{PyString}, x::Py) = pyconvert_return(PyString(x))
13+
14+
Base.ncodeunits(x::PyString) = x.length
15+
Base.codeunit(::PyString) = UInt8
16+
Base.@propagate_inbounds function Base.codeunit(x::PyString, i::Integer)
17+
@boundscheck checkbounds(1:x.length, i)
18+
return unsafe_load(x.ptr, i)
19+
end
20+
Base.isvalid(x::PyString, i::Int) = Utils.utf8_isvalid(x, x.length, i)
21+
Base.iterate(x::PyString, i::Int = 1) = Utils.utf8_iterate(x, x.length, i)

‎src/Wrap/Wrap.jl‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@ using ..Convert
1414
using ..PyMacro
1515

1616
import ..PythonCall:
17-
PyArray, PyDict, PyIO, PyIterable, PyList, PyPandasDataFrame, PySet, PyTable
17+
PyArray, PyDict, PyIO, PyIterable, PyList, PyPandasDataFrame, PySet, PyString, PyTable
1818

1919
using Base: @propagate_inbounds
2020
using Tables: Tables
@@ -26,6 +26,7 @@ include("PyIterable.jl")
2626
include("PyDict.jl")
2727
include("PyList.jl")
2828
include("PySet.jl")
29+
include("PyString.jl")
2930
include("PyArray.jl")
3031
include("PyIO.jl")
3132
include("PyTable.jl")
@@ -73,6 +74,7 @@ function __init__()
7374
end
7475

7576
priority = PYCONVERT_PRIORITY_NORMAL
77+
pyconvert_add_rule("builtins:str", PyString, pyconvert_rule_string, priority)
7678
pyconvert_add_rule("<arraystruct>", Array, pyconvert_rule_array, priority)
7779
pyconvert_add_rule("<arrayinterface>", Array, pyconvert_rule_array, priority)
7880
pyconvert_add_rule("<array>", Array, pyconvert_rule_array, priority)

‎test/Convert.jl‎

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -93,6 +93,28 @@ end
9393
@test x2 === "αβγℵ√"
9494
end
9595

96+
@testitem "str → PyString" begin
97+
x = pystr("aβℵ\0🙂")
98+
s = pyconvert(PyString, x)
99+
@test s isa PyString
100+
@test Py(s) === s.py
101+
@test String(s) == "aβℵ\0🙂"
102+
@test ncodeunits(s) == 11
103+
@test codeunit(s) === UInt8
104+
@test collect(s) == ['a', 'β', 'ℵ', '\0', '🙂']
105+
@test isvalid(s, 1)
106+
@test !isvalid(s, 3)
107+
@test codeunit(s, 3) == 0xb2
108+
@test_throws StringIndexError s[3]
109+
@test_throws BoundsError codeunit(s, 12)
110+
111+
t = PyString("hello")
112+
@test String(t) == "hello"
113+
@test isempty(PyString(""))
114+
@test_throws PyException PyString(pyint(1))
115+
@test pyconvert(Any, x) isa String
116+
end
117+
96118
@testitem "str → Symbol" begin
97119
x1 = pyconvert(Symbol, pystr("hello"))
98120
@test x1 === :hello

‎test/Utils.jl‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,3 +22,12 @@ end
2222
@test s[1:2] == "ab"
2323
@test s[1:2:end] == "aaaab"
2424
end
25+
26+
@testitem "StaticString UTF-8" begin
27+
S = PythonCall.Utils.StaticString
28+
s = S{UInt8,10}("aβℵ🙂")
29+
@test String(s) == "aβℵ🙂"
30+
@test collect(s) == ['a', 'β', 'ℵ', '🙂']
31+
@test !isvalid(S{UInt8,1}((0x80,)), 1)
32+
@test !isvalid(S{UInt8,1}((0xf8,)), 1)
33+
end

0 commit comments

Comments
 (0)