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
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,38 @@
static PyObject* context_copy_current(PyObject* unused, PyObject* args) {
return PyContext_CopyCurrent();
}
static PyObject* contextvar_new(PyObject* unused, PyObject* args) {
const char *name;
PyObject *def = NULL;
if (!PyArg_ParseTuple(args, "s|O", &name, &def))
return NULL;
return PyContextVar_New(name, def);
}
static PyObject* contextvar_get(PyObject* unused, PyObject* args) {
PyObject *var, *value, *def = NULL;
if (!PyArg_ParseTuple(args, "O|O", &var, &def))
return NULL;
if (PyContextVar_Get(var, def, &value) < 0)
return NULL;
if (value == NULL)
Py_RETURN_NONE;
return value;
}
static PyObject* contextvar_set(PyObject* unused, PyObject* args) {
PyObject *var, *value;
if (!PyArg_ParseTuple(args, "OO", &var, &value))
return NULL;
return PyContextVar_Set(var, value);
}
static PyObject* contextvar_reset(PyObject* unused, PyObject* args) {
PyObject *var, *token;
if (!PyArg_ParseTuple(args, "OO", &var, &token))
return NULL;
int result = PyContextVar_Reset(var, token);
if (result < 0)
return NULL;
return PyLong_FromLong(result);
}
static PyObject* context_is_exact_type(PyObject* unused, PyObject* args) {
PyObject *obj;
int kind;
Expand All @@ -85,6 +117,10 @@
}
''',
tp_methods='''
{"var_new", (PyCFunction)contextvar_new, METH_VARARGS | METH_STATIC, ""},
{"var_get", (PyCFunction)contextvar_get, METH_VARARGS | METH_STATIC, ""},
{"var_set", (PyCFunction)contextvar_set, METH_VARARGS | METH_STATIC, ""},
{"var_reset", (PyCFunction)contextvar_reset, METH_VARARGS | METH_STATIC, ""},
{"enter", (PyCFunction)context_enter, METH_VARARGS | METH_STATIC, ""},
{"exit", (PyCFunction)context_exit, METH_VARARGS | METH_STATIC, ""},
{"copy", (PyCFunction)context_copy, METH_VARARGS | METH_STATIC, ""},
Expand Down Expand Up @@ -138,3 +174,59 @@ def test_cext_context_management():

new_ctx = ContextHelper.new()
assert new_ctx.run(v.get) == 'default value'



def test_cext_contextvar_reset():
for default_args in ((), ('default value',), (None,)):
var = ContextHelper.var_new('test_reset', *default_args)
assert isinstance(var, contextvars.ContextVar)
assert ContextHelper.var_get(var) == (default_args[0] if default_args else None)
assert ContextHelper.var_get(var, 'fallback') == 'fallback'

token = ContextHelper.var_set(var, 'first value')
assert token.old_value is contextvars.Token.MISSING
assert ContextHelper.var_get(var) == 'first value'
inner_token = ContextHelper.var_set(var, None)
assert inner_token.old_value == 'first value'
assert ContextHelper.var_reset(var, inner_token) == 0
assert var.get() == 'first value'
assert ContextHelper.var_reset(var, token) == 0
assert var not in contextvars.copy_context()
if default_args:
assert var.get() == default_args[0]
else:
assert_raises(LookupError, var.get)

# Tokens created in Python must work with the C API and vice versa.
token = var.set('python value')
assert ContextHelper.var_reset(var, token) == 0
token = ContextHelper.var_set(var, 'native value')
var.reset(token)
assert var not in contextvars.copy_context()

token = ContextHelper.var_set(var, None)
inner_token = ContextHelper.var_set(var, 'new value')
assert ContextHelper.var_reset(var, inner_token) == 0
assert var.get() is None
assert ContextHelper.var_reset(var, token) == 0


def test_cext_contextvar_reset_errors():
var = contextvars.ContextVar('test_reset_errors')
other_var = contextvars.ContextVar('other_var')
token = ContextHelper.var_set(var, 'value')
assert_raises(TypeError, ContextHelper.var_reset, object(), token, err_check='instance of ContextVar')
assert_raises(TypeError, ContextHelper.var_reset, var, object(), err_check='instance of Token')
assert_raises(ValueError, ContextHelper.var_reset, other_var, token, err_check='different ContextVar')
other_context = ContextHelper.copy_current()
assert_raises(ValueError, other_context.run, ContextHelper.var_reset, var, token, err_check='different Context')
assert_raises(ValueError, other_context.run, var.reset, token, err_check='different Context')
assert var.get() == 'value'
assert other_context.run(var.get) == 'value'

# Rejected resets must leave the token usable in its original context.
assert ContextHelper.var_reset(var, token) == 0
assert_raises(RuntimeError, ContextHelper.var_reset, var, token, err_check='already been used once')
assert_raises(RuntimeError, ContextHelper.var_reset, other_var, token, err_check='already been used once')
assert var not in contextvars.copy_context()
Original file line number Diff line number Diff line change
Expand Up @@ -47,8 +47,8 @@
import static com.oracle.graal.python.builtins.objects.cext.capi.transitions.ArgDescriptor.Pointer;
import static com.oracle.graal.python.builtins.objects.cext.capi.transitions.ArgDescriptor.PyObjectRawPointer;
import static com.oracle.graal.python.builtins.objects.cext.capi.transitions.ArgDescriptor.VoidNoReturn;
import static com.oracle.graal.python.runtime.nativeaccess.NativeMemory.NULLPTR;
import static com.oracle.graal.python.runtime.exception.ExceptionUtils.printPythonLikeStackTrace;
import static com.oracle.graal.python.runtime.nativeaccess.NativeMemory.NULLPTR;

import com.oracle.graal.python.PythonLanguage;
import com.oracle.graal.python.builtins.PythonBuiltinClassType;
Expand All @@ -60,6 +60,7 @@
import com.oracle.graal.python.builtins.objects.cext.capi.transitions.CApiTransitions.PythonToNativeInternalNode;
import com.oracle.graal.python.builtins.objects.contextvars.PContextVar;
import com.oracle.graal.python.builtins.objects.contextvars.PContextVarsContext;
import com.oracle.graal.python.builtins.objects.contextvars.PContextVarsToken;
import com.oracle.graal.python.lib.PyContextCopyCurrent;
import com.oracle.graal.python.nodes.ErrorMessages;
import com.oracle.graal.python.nodes.PRaiseNode;
Expand Down Expand Up @@ -118,7 +119,23 @@ static long PyContextVar_Set(long varPtr, long valPtr) {
PythonContext.PythonThreadState threadState = pythonContext.getThreadState(language);
Object oldValue = pvar.getValue(null, threadState);
pvar.setValue(null, threadState, val);
return PythonToNativeInternalNode.executeNewRefUncached(PFactory.createContextVarsToken(language, pvar, oldValue));
return PythonToNativeInternalNode.executeNewRefUncached(PFactory.createContextVarsToken(language, pvar, threadState.getContextVarsContext(null), oldValue));
}

@CApiBuiltin(ret = Int, args = {PyObjectRawPointer, PyObjectRawPointer}, call = Direct)
static int PyContextVar_Reset(long varPtr, long tokenPtr) {
Object var = NativeToPythonInternalNode.executeUncached(varPtr, false);
if (!(var instanceof PContextVar pvar)) {
throw PRaiseNode.raiseStatic(null, PythonBuiltinClassType.TypeError, ErrorMessages.INSTANCE_OF_CONTEXTVAR_EXPECTED);
}
Object token = NativeToPythonInternalNode.executeUncached(tokenPtr, false);
if (!(token instanceof PContextVarsToken ptoken)) {
throw PRaiseNode.raiseStatic(null, PythonBuiltinClassType.TypeError, ErrorMessages.INSTANCE_OF_TOKEN_EXPECTED, token);
}
PythonContext pythonContext = PythonContext.get(null);
PythonContext.PythonThreadState threadState = pythonContext.getThreadState(pythonContext.getLanguage());
pvar.resetValue(null, threadState, ptoken);
return 0;
}

@CApiBuiltin(ret = PyObjectRawPointer, call = Direct)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -697,7 +697,6 @@ public final class CApiFunction {
@CApiBuiltin(name = "PyConfig_SetBytesString", ret = PYSTATUS, args = {PYCONFIG_PTR, WCHAR_T_PTR_LIST, ConstCharPtr}, call = NotImplemented)
@CApiBuiltin(name = "PyConfig_SetString", ret = PYSTATUS, args = {PYCONFIG_PTR, WCHAR_T_PTR_LIST, CONST_WCHAR_PTR}, call = NotImplemented)
@CApiBuiltin(name = "PyConfig_SetWideStringList", ret = PYSTATUS, args = {PYCONFIG_PTR, PYWIDESTRINGLIST_PTR, Py_ssize_t, WCHAR_T_PTR_LIST}, call = NotImplemented)
@CApiBuiltin(name = "PyContextVar_Reset", ret = Int, args = {PyObject, PyObject}, call = NotImplemented)
@CApiBuiltin(name = "PyCoro_New", ret = PyObject, args = {PyFrameObject, PyObject, PyObject}, call = NotImplemented)
@CApiBuiltin(name = "PyCriticalSection2_Begin", ret = Void, args = {PY_CRITICAL_SECTION2_PTR, PyObjectReturn, PyObjectReturn}, call = NotImplemented)
@CApiBuiltin(name = "PyCriticalSection2_End", ret = Void, args = {PY_CRITICAL_SECTION2_PTR}, call = NotImplemented)
Expand Down
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
/*
* Copyright (c) 2022, 2025, Oracle and/or its affiliates. All rights reserved.
* Copyright (c) 2022, 2026, Oracle and/or its affiliates. All rights reserved.
* DO NOT ALTER OR REMOVE COPYRIGHT NOTICES OR THIS FILE HEADER.
*
* The Universal Permissive License (UPL), Version 1.0
Expand Down Expand Up @@ -42,7 +42,6 @@

import static com.oracle.graal.python.builtins.PythonBuiltinClassType.LookupError;
import static com.oracle.graal.python.builtins.PythonBuiltinClassType.TypeError;
import static com.oracle.graal.python.builtins.PythonBuiltinClassType.ValueError;
import static com.oracle.graal.python.nodes.PGuards.isNoValue;
import static com.oracle.graal.python.nodes.SpecialMethodNames.J___CLASS_GETITEM__;

Expand Down Expand Up @@ -72,7 +71,6 @@
import com.oracle.graal.python.runtime.object.PFactory;
import com.oracle.truffle.api.dsl.Bind;
import com.oracle.truffle.api.dsl.Cached;
import com.oracle.truffle.api.dsl.Cached.Shared;
import com.oracle.truffle.api.dsl.GenerateNodeFactory;
import com.oracle.truffle.api.dsl.NodeFactory;
import com.oracle.truffle.api.dsl.Specialization;
Expand Down Expand Up @@ -150,7 +148,7 @@ static Object set(VirtualFrame frame, PContextVar self, Object value,
} finally {
BoundaryCallContext.exit(frame, boundaryCallData, saved);
}
return PFactory.createContextVarsToken(language, self, oldValue);
return PFactory.createContextVarsToken(language, self, threadState.getContextVarsContext(inliningTarget), oldValue);
}
}

Expand All @@ -161,33 +159,21 @@ public abstract static class ResetNode extends PythonBinaryBuiltinNode {
static Object reset(VirtualFrame frame, PContextVar self, PContextVarsToken token,
@Bind Node inliningTarget,
@Bind PythonContext pythonContext,
@Cached("createFor($node)") BoundaryCallData boundaryCallData,
@Shared @Cached PRaiseNode raise) {
if (self == token.getVar()) {
token.use(inliningTarget, raise);
PythonContext.PythonThreadState threadState = pythonContext.getThreadState(pythonContext.getLanguage(inliningTarget));
Object saved = BoundaryCallContext.enter(frame, boundaryCallData);
try {
if (token.getOldValue() == null) {
PContextVarsContext context = threadState.getContextVarsContext(inliningTarget);
context.contextVarValues = context.contextVarValues.without(self, self.getHash());
} else {
self.setValue(inliningTarget, threadState, token.getOldValue());
}
} finally {
BoundaryCallContext.exit(frame, boundaryCallData, saved);
}
} else {
throw raise.raise(inliningTarget, ValueError, ErrorMessages.TOKEN_FOR_DIFFERENT_CONTEXTVAR, token);
@Cached("createFor($node)") BoundaryCallData boundaryCallData) {
PythonContext.PythonThreadState threadState = pythonContext.getThreadState(pythonContext.getLanguage(inliningTarget));
Object saved = BoundaryCallContext.enter(frame, boundaryCallData);
try {
self.resetValue(inliningTarget, threadState, token);
} finally {
BoundaryCallContext.exit(frame, boundaryCallData, saved);
}
return PNone.NONE;
}

@Specialization(guards = "!isToken(token)")
Object doError(@SuppressWarnings("unused") PContextVar self, Object token,
@Bind Node inliningTarget,
@Shared @Cached PRaiseNode raise) {
throw raise.raise(inliningTarget, TypeError, ErrorMessages.INSTANCE_OF_TOKEN_EXPECTED, token);
@Bind Node inliningTarget) {
throw PRaiseNode.raiseStatic(inliningTarget, TypeError, ErrorMessages.INSTANCE_OF_TOKEN_EXPECTED, token);
}

static boolean isToken(Object obj) {
Expand Down
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
/*
* Copyright (c) 2022, 2025, Oracle and/or its affiliates. All rights reserved.
* Copyright (c) 2022, 2026, Oracle and/or its affiliates. All rights reserved.
* DO NOT ALTER OR REMOVE COPYRIGHT NOTICES OR THIS FILE HEADER.
*
* The Universal Permissive License (UPL), Version 1.0
Expand Down Expand Up @@ -81,6 +81,16 @@ public void setValue(Node node, PythonContext.PythonThreadState state, Object va
current.contextVarValues = current.contextVarValues.withEntry(new Hamt.Entry(this, getHash(), value));
}

public void resetValue(Node node, PythonContext.PythonThreadState state, PContextVarsToken token) {
PContextVarsContext current = state.getContextVarsContext(node);
token.use(node, this, current);
if (token.getOldValue() == null) {
current.contextVarValues = current.contextVarValues.without(this, getHash());
} else {
setValue(node, state, token.getOldValue());
}
}

public Object get(Node node, PythonContext.PythonThreadState state, Object defaultValue) {
Object result = getValue(node, state);
if (result != null) {
Expand Down
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
/*
* Copyright (c) 2022, 2025, Oracle and/or its affiliates. All rights reserved.
* Copyright (c) 2022, 2026, Oracle and/or its affiliates. All rights reserved.
* DO NOT ALTER OR REMOVE COPYRIGHT NOTICES OR THIS FILE HEADER.
*
* The Universal Permissive License (UPL), Version 1.0
Expand Down Expand Up @@ -50,19 +50,27 @@
public class PContextVarsToken extends PythonBuiltinObject {
public static final Object MISSING = new Object();
private final PContextVar var;
private final PContextVarsContext context;
private final Object oldValue;

private boolean used = false;

public PContextVarsToken(PContextVar var, Object oldValue, Object cls, Shape instanceShape) {
public PContextVarsToken(PContextVar var, PContextVarsContext context, Object oldValue, Object cls, Shape instanceShape) {
super(cls, instanceShape);
this.var = var;
this.context = context;
this.oldValue = oldValue;
}

public void use(Node inliningTarget, PRaiseNode raise) {
public void use(Node inliningTarget, PContextVar expectedVar, PContextVarsContext expectedContext) {
if (used) {
throw raise.raise(inliningTarget, PythonBuiltinClassType.RuntimeError, ErrorMessages.TOKEN_ALREADY_USED, this);
throw PRaiseNode.raiseStatic(inliningTarget, PythonBuiltinClassType.RuntimeError, ErrorMessages.TOKEN_ALREADY_USED, this);
}
if (var != expectedVar) {
throw PRaiseNode.raiseStatic(inliningTarget, PythonBuiltinClassType.ValueError, ErrorMessages.TOKEN_FOR_DIFFERENT_CONTEXTVAR, this);
}
if (context != expectedContext) {
throw PRaiseNode.raiseStatic(inliningTarget, PythonBuiltinClassType.ValueError, ErrorMessages.TOKEN_FOR_DIFFERENT_CONTEXT, this);
}
used = true;
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1379,6 +1379,7 @@ public abstract class ErrorMessages {
public static final TruffleString CONTEXTVAR_KEY_EXPECTED = tsLiteral("A ContextVar key was expected, got %s");
public static final TruffleString CANNOT_ENTER_CONTEXT_ALREADY_ENTERED = tsLiteral("cannot enter context: %s is already entered");
public static final TruffleString TOKEN_ONLY_BY_CONTEXTVAR = tsLiteral("Tokens can only be created by ContextVars");
public static final TruffleString TOKEN_FOR_DIFFERENT_CONTEXT = tsLiteral("%s was created in a different Context");
public static final TruffleString TOKEN_FOR_DIFFERENT_CONTEXTVAR = tsLiteral("%s was created by a different ContextVar");

public static final TruffleString ATTRIBUTE_VALUE_MUST_BE_BOOL = tsLiteral("attribute value type must be bool");
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1498,8 +1498,8 @@ public static PContextVarsContext copyContextVarsContext(PythonLanguage language
return new PContextVarsContext(original, PythonBuiltinClassType.ContextVarsContext, PythonBuiltinClassType.ContextVarsContext.getInstanceShape(language));
}

public static PContextVarsToken createContextVarsToken(PythonLanguage language, PContextVar var, Object oldValue) {
return new PContextVarsToken(var, oldValue, PythonBuiltinClassType.ContextVarsToken, PythonBuiltinClassType.ContextVarsToken.getInstanceShape(language));
public static PContextVarsToken createContextVarsToken(PythonLanguage language, PContextVar var, PContextVarsContext context, Object oldValue) {
return new PContextVarsToken(var, context, oldValue, PythonBuiltinClassType.ContextVarsToken, PythonBuiltinClassType.ContextVarsToken.getInstanceShape(language));
}

public static PGenericAlias createGenericAlias(PythonLanguage language, Object cls, Shape shape, Object origin, PTuple arguments, boolean starred) {
Expand Down
Loading