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
109 changes: 96 additions & 13 deletions src/C/regions.jl
Original file line number Diff line number Diff line change
@@ -1,27 +1,83 @@
const PyThreadStatePtr = Ptr{Cvoid}

mutable struct ThreadState
sem::Base.Semaphore
tstate::PyThreadStatePtr
end
ThreadState() = ThreadState(Base.Semaphore(1), C_NULL)

mutable struct TaskState
# CPython state used by this task for the duration of its outermost region.
tstate::PyThreadStatePtr
# Semaphore protecting `tstate`; `nothing` when the task is outside a session.
sem::Union{Nothing,Base.Semaphore}
# Whether `tstate` is currently attached to the task's pinned OS thread.
attached::Bool
# Julia thread to which the task is pinned while its session is active.
tid::Int
# Stickiness to restore when the task's outermost region or break finishes.
oldsticky::Bool
end
TaskState() = TaskState(C_NULL, nothing, false, 0, false)

mutable struct ThreadState
# Serializes tasks which use the persistent CPython state on this Julia thread.
sem::Base.Semaphore
# PythonCall-owned CPython state for this Julia thread, allocated lazily.
tstate::PyThreadStatePtr
# Task which currently holds `sem` with a CPython state attached, if any.
task::Union{Nothing,Task}
# State belonging to `task`; cached so nested regions need no task-local lookup.
task_state::Union{Nothing,TaskState}
end
ThreadState() = ThreadState(Base.Semaphore(1), C_NULL, nothing, nothing)

const THREAD_STATE = Utils.OncePerThread{ThreadState}(ThreadState)
const TASK_STATE = Utils.OncePerTask{TaskState}(TaskState)

# A region session starts at the outermost region entered by a Julia task and ends
# when that region exits. The task is made sticky for the whole session: CPython
# thread states must be restored on the same OS thread on which they were saved.
# The corresponding ThreadState semaphore gives the task exclusive use of that
# thread's persistent PyThreadState. A Python-originated call already has a CPython
# state attached; that state is borrowed instead, but is protected by the same
# semaphore so PythonCall's accounting remains identical.
#
# Nested regions merely return `:noop`. A break saves/detaches the CPython state,
# clears the thread cache, and releases the semaphore, allowing other sticky tasks
# to use this Julia thread while the original task yields. On return it reacquires
# the semaphore before restoring both the CPython state and cache. Consequently the
# cache is populated exactly while a task owns the semaphore with Python attached.
# This invariant is what makes its lock-free fast path safe.
#
# An attached CPython state is also a still cheaper fast path for @pyregion itself.
# Given that PythonCall is the only Julia-side entry to Python, an attached state
# means either that an enclosing region already did the bookkeeping or that Python
# called into Julia and lent us its state. In either case a nested region has no
# transition to perform. Conversely, a region inside a break observes no attached
# state and takes the full path below, reacquiring and restoring its task state.

current_tstate() = PyThreadState_GetUnchecked()
has_tstate() = CTX.is_initialized && current_tstate() != C_NULL

# An attached task is pinned to this Julia thread and owns its semaphore, so the
# thread-local cache cannot change underneath it. This makes nested regions avoid
# the considerably more expensive OncePerTask lookup. The cache is cleared before
# releasing the semaphore and restored after reacquiring it around a region break.
function current_task_state(task::Task)
ts = THREAD_STATE()
return ts.task === task ? ts.task_state::TaskState : TASK_STATE()
end

function set_task!(ts::ThreadState, task::Task, s::TaskState)
ts.task = task
ts.task_state = s
return
end

function clear_task!(ts::ThreadState)
ts.task = nothing
ts.task_state = nothing
return
end

function reset!(s::TaskState)
# `oldsticky` is deliberately left alone: start_session! overwrites it before
# the TaskState becomes observable as an active session again.
s.tstate = C_NULL
s.sem = nothing
s.attached = false
Expand All @@ -30,6 +86,8 @@ function reset!(s::TaskState)
end

function start_session!(task::Task, s::TaskState)
# Pin before recording the thread and looking up its ThreadState. Once sticky,
# cooperative scheduling cannot resume this task on another Julia thread.
s.oldsticky = task.sticky
task.sticky = true
s.tid = Threads.threadid()
Expand All @@ -45,6 +103,8 @@ end

function enter_region(task::Task, s::TaskState)
if s.tstate != C_NULL
# An existing session is either already attached (ordinary nesting), or
# temporarily detached because this region is nested inside a break.
check_thread(task, s)
if s.attached
return :noop
Expand All @@ -53,6 +113,7 @@ function enter_region(task::Task, s::TaskState)
try
PyEval_RestoreThread(s.tstate)
s.attached = true
set_task!(THREAD_STATE(), task, s)
catch
Base.release(s.sem::Base.Semaphore)
rethrow()
Expand All @@ -73,8 +134,10 @@ function enter_region(task::Task, s::TaskState)
acquired = true
current = current_tstate()
if current != C_NULL
# Calls originating in Python must return with its state still attached.
s.tstate = current
s.attached = true
set_task!(ts, task, s)
return :root_borrowed
end
if ts.tstate == C_NULL
Expand All @@ -84,6 +147,7 @@ function enter_region(task::Task, s::TaskState)
s.tstate = ts.tstate
PyEval_RestoreThread(s.tstate)
s.attached = true
set_task!(ts, task, s)
return :root_owned
catch
acquired && Base.release(ts.sem)
Expand All @@ -94,23 +158,28 @@ function enter_region(task::Task, s::TaskState)
end

function exit_region(task::Task, s::TaskState, token)
# Tokens encode precisely which work enter_region performed, so every path
# reverses only its own transition and nested :noop regions do no attach work.
token === :noop && return
check_thread(task, s)
if token === :detach
saved = PyEval_SaveThread()
@assert saved == s.tstate
s.attached = false
clear_task!(THREAD_STATE())
Base.release(s.sem::Base.Semaphore)
elseif token === :root_owned
saved = PyEval_SaveThread()
@assert saved == s.tstate
sem, oldsticky = s.sem::Base.Semaphore, s.oldsticky
clear_task!(THREAD_STATE())
reset!(s)
Base.release(sem)
task.sticky = oldsticky
elseif token === :root_borrowed
@assert current_tstate() == s.tstate
sem, oldsticky = s.sem::Base.Semaphore, s.oldsticky
clear_task!(THREAD_STATE())
reset!(s)
Base.release(sem)
task.sticky = oldsticky
Expand All @@ -122,14 +191,18 @@ end

function enter_break(task::Task, s::TaskState)
if s.tstate != C_NULL
# Only the outermost break detaches. Further breaks while detached are no-ops.
check_thread(task, s)
!s.attached && return :noop
saved = PyEval_SaveThread()
@assert saved == s.tstate
s.attached = false
clear_task!(THREAD_STATE())
Base.release(s.sem::Base.Semaphore)
return :restore
end
# Outside a region, a break matters only in a Python-originated callback. It
# borrows and detaches that state so yielding Julia code cannot retain Python.
current_tstate() == C_NULL && return :noop

ts = try
Expand Down Expand Up @@ -170,8 +243,12 @@ function exit_break(task::Task, s::TaskState, token)
Base.acquire(s.sem::Base.Semaphore)
PyEval_RestoreThread(s.tstate)
s.attached = true
set_task!(THREAD_STATE(), task, s)
if token === :root_restore
# The borrowed Python state stays attached for the caller, but the temporary
# Julia session and its cache ownership end here.
sem, oldsticky = s.sem::Base.Semaphore, s.oldsticky
clear_task!(THREAD_STATE())
reset!(s)
Base.release(sem)
task.sticky = oldsticky
Expand All @@ -191,13 +268,19 @@ blocks cooperatively in [`@pyregionbreak`](@ref).
"""
macro pyregion(ex)
quote
local task = current_task()
local state = $TASK_STATE()
local token = $enter_region(task, state)
try
if $current_tstate() != C_NULL
# The CPython TLS lookup is enough to prove that this region is nested
# or Python-originated; avoid both Julia task- and thread-local lookups.
$(esc(ex))
finally
$exit_region(task, state, token)
else
local task = current_task()
local state = $current_task_state(task)
local token = $enter_region(task, state)
try
$(esc(ex))
finally
$exit_region(task, state, token)
end
end
end
end
Expand All @@ -212,7 +295,7 @@ PythonCall operations still work automatically, and both kinds of region nest fr
macro pyregionbreak(ex)
quote
local task = current_task()
local state = $TASK_STATE()
local state = $current_task_state(task)
local token = $enter_break(task, state)
try
$(esc(ex))
Expand Down
58 changes: 58 additions & 0 deletions test/C.jl
Original file line number Diff line number Diff line change
Expand Up @@ -16,3 +16,61 @@
@test PythonCall.python_version().major == 3
end
end

@testitem "Python regions and task-safe thread states" setup = [Setup] begin
using Base.Threads

@test PythonCall.C.PyThreadState_GetUnchecked() == C_NULL
@test @pyregion pyconvert(Int, pyint(12)) == 12
@test PythonCall.C.PyThreadState_GetUnchecked() == C_NULL

# Nested regions reuse the state cached by the owning Julia thread. A break
# must relinquish that cache while yielding and restore it on return.
@pyregion begin
local state = PythonCall.C.TASK_STATE()
local thread_state = PythonCall.C.THREAD_STATE()
@test thread_state.task === current_task()
@test thread_state.task_state === state
@pyregion @test PythonCall.C.current_task_state(current_task()) === state
@pyregionbreak begin
@test thread_state.task === nothing
yield()
end
@test thread_state.task === current_task()
@test thread_state.task_state === state
end
@test PythonCall.C.THREAD_STATE().task === nothing

@test @pyregion begin
@pyregion pyconvert(Int, pyint(1)) == 1
@pyregionbreak begin
yield()
@pyregion pyconvert(Int, pyint(2)) == 2
@pyregionbreak yield()
end
pyconvert(Int, pyint(3)) == 3
end

@test_throws ErrorException @pyregion @pyregionbreak error("region exception")
@test PythonCall.C.PyThreadState_GetUnchecked() == C_NULL
@test_throws PyException @pyregion pybuiltins.int("not an integer")
@test pyconvert(Int, pyint(5)) == 5

results = fetch.([@spawn begin
total = 0
for i in 1:100
total += pyconvert(Int, pyint(i))
@pyregionbreak yield()
end
total
end for _ in 1:max(8, 2nthreads())])
@test all(==(5050), results)

f = pyfunc() do
yield()
pyconvert(Int, pybuiltins.sum([1, 2, 3]))
end
@test pyconvert(Int, f()) == 6

@test pyconvert(Int, @py 1 + @jl(pyconvert(Int, pyint(2)))) == 3
end
42 changes: 0 additions & 42 deletions test/Region.jl

This file was deleted.

Loading