Compare commits

...
6 Commits
Author SHA1 Message Date
rcarrigaandItai Bohadana fcf211872a docs: remove nvim-treesitter ref 2026-08-18 16:40:54 +03:00
rcarrigaandItai Bohadana 406fd3db4a feat: remove plenary 2026-08-18 16:40:52 +03:00
f48dbfa52b chore: Replace use of deprecated vim.tbl_flatten (#112)
vim.tbl_flatten is deprecated and will be removed in 0.13. Just replaced
it with the recommendation in the deprecation notice.

For some reason discovering tests did not work anymore after upgrading from
0.11.x to 0.12.1 even though vim.tbl_flatten still works in this
version. Replacing it fixes test discovery for me though and it needs to
be done before 0.13 anyway.

Co-authored-by: Sam Castelain <sam@secury-360.com>
2026-08-18 16:40:31 +03:00
Thomas VandalandItai Bohadana 460afdf404 fix(pytest): replace rootpath with rootdir (#106)
* Replace rootdir with rootpath

* Add try/except to handle older pytest versions as well

* Extract try/except to a function and use in `pytest_runtest_makereport`
2026-08-18 16:39:59 +03:00
Salomon PoppandItai Bohadana b720009f3e fix: propagate exit code (#108)
* fix(pytest): propagate exit code

* feat: derive exit code for unittest
2026-08-18 16:39:57 +03:00
SpaceShamanandItai Bohadana 7d7cacb91a feat(pytest): support pytest-xdist (#105) 2026-08-18 16:39:08 +03:00
8 changed files with 114 additions and 65 deletions
+1 -1
View File
@@ -3,7 +3,7 @@
[Neotest](https://github.com/rcarriga/neotest) adapter for python. [Neotest](https://github.com/rcarriga/neotest) adapter for python.
Supports Pytest and unittest test files. Supports Pytest and unittest test files.
Requires [nvim-treesitter](https://github.com/nvim-treesitter/nvim-treesitter) and the parser for python. Requires the treesitter parser for python.
```lua ```lua
require("neotest").setup({ require("neotest").setup({
+31 -13
View File
@@ -1,6 +1,5 @@
local nio = require("nio") local nio = require("nio")
local lib = require("neotest.lib") local lib = require("neotest.lib")
local Path = require("plenary.path")
local M = {} local M = {}
@@ -8,7 +7,7 @@ function M.is_test_file(file_path)
if not vim.endswith(file_path, ".py") then if not vim.endswith(file_path, ".py") then
return false return false
end end
local elems = vim.split(file_path, Path.path.sep) local elems = vim.split(file_path, lib.files.sep)
local file_name = elems[#elems] local file_name = elems[#elems]
return vim.startswith(file_name, "test_") or vim.endswith(file_name, "_test.py") return vim.startswith(file_name, "test_") or vim.endswith(file_name, "_test.py")
end end
@@ -35,14 +34,14 @@ function M.get_python_command(root)
end end
-- Use activated virtualenv. -- Use activated virtualenv.
if vim.env.VIRTUAL_ENV then if vim.env.VIRTUAL_ENV then
python_command_mem[root] = { Path:new(vim.env.VIRTUAL_ENV, venv_bin, "python").filename } python_command_mem[root] = { vim.fs.joinpath(vim.env.VIRTUAL_ENV, venv_bin, "python") }
return python_command_mem[root] return python_command_mem[root]
end end
for _, pattern in ipairs({ "*", ".*" }) do for _, pattern in ipairs({ "*", ".*" }) do
local match = nio.fn.glob(Path:new(root or nio.fn.getcwd(), pattern, "pyvenv.cfg").filename) local match = nio.fn.glob(vim.fs.joinpath(root or nio.fn.getcwd(), pattern, "pyvenv.cfg"))
if match ~= "" then if match ~= "" then
python_command_mem[root] = { (Path:new(match):parent() / venv_bin / "python").filename } python_command_mem[root] = { vim.fs.joinpath(vim.fs.dirname(match), venv_bin, "python") }
return python_command_mem[root] return python_command_mem[root]
end end
end end
@@ -52,7 +51,7 @@ function M.get_python_command(root)
if success and exit_code == 0 then if success and exit_code == 0 then
local venv = data.stdout:gsub("\r?\n", "") local venv = data.stdout:gsub("\r?\n", "")
if venv then if venv then
python_command_mem[root] = { Path:new(venv).filename } python_command_mem[root] = { venv }
return python_command_mem[root] return python_command_mem[root]
end end
end end
@@ -67,7 +66,7 @@ function M.get_python_command(root)
if success and exit_code == 0 then if success and exit_code == 0 then
local venv = data.stdout:gsub("\r?\n", "") local venv = data.stdout:gsub("\r?\n", "")
if venv then if venv then
python_command_mem[root] = { Path:new(venv, venv_bin, "python").filename } python_command_mem[root] = { vim.fs.joinpath(venv, venv_bin, "python") }
return python_command_mem[root] return python_command_mem[root]
end end
end end
@@ -80,7 +79,7 @@ function M.get_python_command(root)
{ stdout = true } { stdout = true }
) )
if success and exit_code == 0 then if success and exit_code == 0 then
python_command_mem[root] = { Path:new(data).filename } python_command_mem[root] = { data }
return python_command_mem[root] return python_command_mem[root]
end end
end end
@@ -110,6 +109,25 @@ end
---@return string ---@return string
local function scan_test_function_pattern(runner, config, python_command) local function scan_test_function_pattern(runner, config, python_command)
local test_function_pattern = "^test" local test_function_pattern = "^test"
if runner == "pytest" and config.pytest_discovery then
<<<<<<< HEAD
local cmd = vim
.iter({ python_command, M.get_script_path(), "--pytest-extract-test-name-template" })
:flatten()
:totable()
=======
local cmd = vim.iter({ python_command, M.get_script_path(), "--pytest-extract-test-name-template" }):flatten()
:totable()
>>>>>>> 51c453d (feat: remove plenary)
local _, data = lib.process.run(cmd, { stdout = true, stderr = true })
for line in vim.gsplit(data.stdout, "\n", true) do
if string.sub(line, 1, 1) == "{" and string.find(line, "python_functions") ~= nil then
local pytest_option = vim.json.decode(line)
test_function_pattern = pytest_option.python_functions
end
end
end
return test_function_pattern return test_function_pattern
end end
@@ -155,7 +173,7 @@ M.treesitter_queries = function(runner, config, python_command)
end end
M.get_root = M.get_root =
lib.files.match_root_pattern("pyproject.toml", "setup.cfg", "mypy.ini", "pytest.ini", "setup.py") lib.files.match_root_pattern("pyproject.toml", "setup.cfg", "mypy.ini", "pytest.ini", "setup.py")
function M.create_dap_config(python_path, script_path, script_args, dap_args) function M.create_dap_config(python_path, script_path, script_args, dap_args)
return vim.tbl_extend("keep", { return vim.tbl_extend("keep", {
@@ -181,13 +199,13 @@ function M.get_runner(python_path)
return "unittest" return "unittest"
end end
if if
vim_test_runner and lib.func_util.index({ "unittest", "pytest", "django" }, vim_test_runner) vim_test_runner and lib.func_util.index({ "unittest", "pytest", "django" }, vim_test_runner)
then then
return vim_test_runner return vim_test_runner
end end
local runner = M.module_exists("pytest_", python_path) and "pytest" local runner = M.module_exists("pytest", python_path) and "pytest"
or M.module_exists("django", python_path) and "django" or M.module_exists("django", python_path) and "django"
or "unittest" or "unittest"
stored_runners[command_str] = runner stored_runners[command_str] = runner
return runner return runner
end end
+1 -1
View File
@@ -17,4 +17,4 @@ with add_to_path():
from neotest_python import main from neotest_python import main
if __name__ == "__main__": if __name__ == "__main__":
main(sys.argv[1:]) sys.exit(main(sys.argv[1:]))
+6 -6
View File
@@ -50,20 +50,18 @@ parser.add_argument(
parser.add_argument("args", nargs="*") parser.add_argument("args", nargs="*")
def main(argv: List[str]): def main(argv: List[str]) -> int:
if "--pytest-collect" in argv: if "--pytest-collect" in argv:
argv.remove("--pytest-collect") argv.remove("--pytest-collect")
from .pytest_ import collect from .pytest_ import collect
collect(argv) return collect(argv)
return
if "--pytest-extract-test-name-template" in argv: if "--pytest-extract-test-name-template" in argv:
argv.remove("--pytest-extract-test-name-template") argv.remove("--pytest-extract-test-name-template")
from .pytest_ import extract_test_name_template from .pytest_ import extract_test_name_template
extract_test_name_template(argv) return extract_test_name_template(argv)
return
args = parser.parse_args(argv) args = parser.parse_args(argv)
adapter = get_adapter(TestRunner(args.runner), args.emit_parameterized_ids) adapter = get_adapter(TestRunner(args.runner), args.emit_parameterized_ids)
@@ -74,7 +72,9 @@ def main(argv: List[str]):
stream_file.write(json.dumps({"id": pos_id, "result": result}) + "\n") stream_file.write(json.dumps({"id": pos_id, "result": result}) + "\n")
stream_file.flush() stream_file.flush()
results = adapter.run(args.args, stream) results, exit_code = adapter.run(args.args, stream)
with open(args.results_file, "w") as results_file: with open(args.results_file, "w") as results_file:
json.dump(results, results_file) json.dump(results, results_file)
return exit_code
+2 -2
View File
@@ -1,6 +1,6 @@
import abc import abc
from enum import Enum from enum import Enum
from typing import TYPE_CHECKING, Callable, Dict, List, Optional from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Tuple
class NeotestResultStatus(str, Enum): class NeotestResultStatus(str, Enum):
@@ -43,6 +43,6 @@ class NeotestAdapter(abc.ABC):
} }
@abc.abstractmethod @abc.abstractmethod
def run(self, args: List[str], stream: Callable): def run(self, args: List[str], stream: Callable) -> Tuple[Dict, int]:
del args, stream del args, stream
raise NotImplementedError raise NotImplementedError
+4 -3
View File
@@ -55,7 +55,7 @@ class DjangoNeotestAdapter(CaseUtilsMixin, NeotestAdapter):
relative_dotted = relative_stem.replace(os.sep, ".") relative_dotted = relative_stem.replace(os.sep, ".")
return [*args, ".".join([relative_dotted, *child_ids])] return [*args, ".".join([relative_dotted, *child_ids])]
def run(self, args: List[str], _) -> Dict: def run(self, args: List[str], _) -> Tuple[Dict, int]:
errs: Dict[str, Tuple[Exception, Any, TracebackType]] = {} errs: Dict[str, Tuple[Exception, Any, TracebackType]] = {}
results = {} results = {}
@@ -146,5 +146,6 @@ class DjangoNeotestAdapter(CaseUtilsMixin, NeotestAdapter):
runner = DjangoUnittestRunner( runner = DjangoUnittestRunner(
**vars(parser.parse_args(argv[1:-1])) # parse plugin config args **vars(parser.parse_args(argv[1:-1])) # parse plugin config args
) )
runner.run_tests(test_labels=[argv[-1]]) # pass test label failures = runner.run_tests(test_labels=[argv[-1]]) # pass test label
return results exit_code = 0 if failures == 0 else 1
return results, exit_code
+65 -35
View File
@@ -1,18 +1,20 @@
from io import StringIO
import json import json
from pathlib import Path
import re import re
from typing import Callable, Dict, List, Optional, Union from typing import Callable, Dict, List, Optional, Union
from . import params_getter from . import params_getter
from io import StringIO
from pathlib import Path
from typing import Callable, Dict, Generator, List, Optional, Tuple, Union
import pytest import pytest
from _pytest._code.code import ExceptionRepr from _pytest._code.code import ExceptionRepr
from _pytest.terminal import TerminalReporter
from _pytest.fixtures import FixtureLookupErrorRepr from _pytest.fixtures import FixtureLookupErrorRepr
from _pytest.terminal import TerminalReporter
from .base import NeotestAdapter, NeotestError, NeotestResult, NeotestResultStatus from .base import NeotestAdapter, NeotestError, NeotestResult, NeotestResultStatus
ANSI_ESCAPE = re.compile(r'\x1B(?:[@-Z\\-_]|\[[0-?]*[ -/]*[@-~])') ANSI_ESCAPE = re.compile(r"\x1B(?:[@-Z\\-_]|\[[0-?]*[ -/]*[@-~])")
class PytestNeotestAdapter(NeotestAdapter): class PytestNeotestAdapter(NeotestAdapter):
def __init__(self, emit_parameterized_ids: bool): def __init__(self, emit_parameterized_ids: bool):
@@ -22,18 +24,18 @@ class PytestNeotestAdapter(NeotestAdapter):
self, self,
args: List[str], args: List[str],
stream: Callable[[str, NeotestResult], None], stream: Callable[[str, NeotestResult], None],
) -> Dict[str, NeotestResult]: ) -> Tuple[Dict[str, NeotestResult], int]:
result_collector = NeotestResultCollector( result_collector = NeotestResultCollector(
self, stream=stream, emit_parameterized_ids=self.emit_parameterized_ids self, stream=stream, emit_parameterized_ids=self.emit_parameterized_ids
) )
pytest.main( exit_code = pytest.main(
args=args, args=args,
plugins=[ plugins=[
result_collector, result_collector,
NeotestDebugpyPlugin(), NeotestDebugpyPlugin(),
], ],
) )
return result_collector.results return result_collector.results, int(exit_code)
class NeotestResultCollector: class NeotestResultCollector:
@@ -72,10 +74,22 @@ class NeotestResultCollector:
buffer.seek(0) buffer.seek(0)
return buffer.read() return buffer.read()
def pytest_configure(self, config: "pytest.Config"):
self.pytest_config = config
def _get_abs_path(self, file_path: Union[str, Path]):
try:
# rootpath is now the preferred way to access root
abs_path = str(self.pytest_config.rootpath / file_path)
except AttributeError:
# fallback to rootdir for older pytest versions
abs_path = str(Path(self.pytest_config.rootdir, file_path))
return abs_path
def pytest_deselected(self, items: List["pytest.Item"]): def pytest_deselected(self, items: List["pytest.Item"]):
for report in items: for report in items:
file_path, *name_path = report.nodeid.split("::") file_path, *name_path = report.nodeid.split("::")
abs_path = str(Path(self.pytest_config.rootdir, file_path)) abs_path = self._get_abs_path(file_path)
*namespaces, test_name = name_path *namespaces, test_name = name_path
valid_test_name, *params = test_name.split("[") # ] valid_test_name, *params = test_name.split("[") # ]
pos_id = "::".join([abs_path, *(namespaces), valid_test_name]) pos_id = "::".join([abs_path, *(namespaces), valid_test_name])
@@ -89,44 +103,61 @@ class NeotestResultCollector:
) )
if not params: if not params:
self.stream(pos_id, result) self.stream(pos_id, result)
self.results[pos_id] = result
def pytest_cmdline_main(self, config: "pytest.Config"): self.results[pos_id] = result
self.pytest_config = config
@pytest.hookimpl(hookwrapper=True) @pytest.hookimpl(hookwrapper=True)
def pytest_runtest_makereport( def pytest_runtest_makereport(
self, item: "pytest.Item", call: "pytest.CallInfo" self, item: "pytest.Item", call: "pytest.CallInfo"
) -> None: ) -> Generator:
# pytest generates the report.outcome field in its internal
# pytest_runtest_makereport implementation, so call it first. (We don't
# implement pytest_runtest_logreport because it doesn't have access to
# call.excinfo.)
outcome = yield outcome = yield
report = outcome.get_result() report: pytest.TestReport = outcome.get_result()
if report.when not in {"call", "setup"} or report.outcome != "failed":
return
exc_repr = report.longrepr
if not isinstance(exc_repr, ExceptionRepr):
return
file_path, *_ = item.nodeid.split("::")
abs_path = str(Path(self.pytest_config.rootdir, file_path))
report.error_line = next(
(
traceback_entry.lineno
for traceback_entry in reversed(call.excinfo.traceback)
if str(traceback_entry.path) == abs_path
),
None,
)
def pytest_runtest_logreport(self, report: "pytest.TestReport") -> None:
if not ( if not (
report.when == "call" report.when == "call"
or (report.when == "setup" and report.outcome in ("skipped", "failed")) or (report.when == "setup" and report.outcome in ("skipped", "failed"))
): ):
return return
file_path, *name_path = item.nodeid.split("::") file_path, *name_path = report.nodeid.split("::")
abs_path = str(Path(self.pytest_config.rootdir, file_path)) abs_path = self._get_abs_path(file_path)
*namespaces, test_name = name_path *namespaces, test_name = name_path
valid_test_name, *params = test_name.split("[") # ] valid_test_name, *params = test_name.split("[") # ]
pos_id = "::".join([abs_path, *namespaces, valid_test_name]) pos_id = "::".join([abs_path, *namespaces, valid_test_name])
errors: List[NeotestError] = [] errors: List[NeotestError] = []
short = self._get_short_output(self.pytest_config, report) short = self._get_short_output(self.pytest_config, report)
msg_prefix = "" msg_prefix = ""
if getattr(item, "callspec", None) is not None: param_id = None
# Parametrized test if "[" in test_name and test_name.endswith("]"):
param_id = test_name[len(valid_test_name) + 1 : -1]
if param_id:
if self.emit_parameterized_ids: if self.emit_parameterized_ids:
pos_id += f"[{item.callspec.id}]" pos_id += f"[{param_id}]"
else: else:
msg_prefix = f"[{item.callspec.id}] " msg_prefix = f"[{param_id}] "
if report.outcome == "failed": if report.outcome == "failed":
exc_repr = report.longrepr exc_repr = report.longrepr
# Test fails due to condition outside of test e.g. xfail # Test fails due to condition outside of test e.g. xfail
@@ -134,27 +165,27 @@ class NeotestResultCollector:
errors.append({"message": msg_prefix + exc_repr, "line": None}) errors.append({"message": msg_prefix + exc_repr, "line": None})
# Test failed internally # Test failed internally
elif isinstance(exc_repr, ExceptionRepr): elif isinstance(exc_repr, ExceptionRepr):
error_message = ANSI_ESCAPE.sub('', exc_repr.reprcrash.message) # type: ignore # Try to use reprcrash, but ensure the line is 0-based
error_line = None error_message = ANSI_ESCAPE.sub("", exc_repr.reprcrash.message) # type: ignore
for traceback_entry in reversed(call.excinfo.traceback): # error_line = report.error_line
if str(traceback_entry.path) == abs_path: error_line = getattr(report, "error_line", None)
error_line = traceback_entry.lineno
errors.append( errors.append(
{"message": msg_prefix + error_message, "line": error_line} {"message": msg_prefix + error_message, "line": error_line}
) )
elif isinstance(exc_repr, FixtureLookupErrorRepr): elif isinstance(exc_repr, FixtureLookupErrorRepr):
line0 = getattr(exc_repr, "firstlineno", None)
if isinstance(line0, int):
line0 = max(0, line0 - 1) # 0-based
errors.append( errors.append(
{ {
"message": msg_prefix + exc_repr.errorstring, "message": msg_prefix + exc_repr.errorstring,
"line": exc_repr.firstlineno, "line": line0,
} }
) )
else: else:
# TODO: Figure out how these are returned and how to represent # Preserve compatibility with previous behavior
raise Exception( errors.append({"message": msg_prefix + str(exc_repr), "line": None})
f"Unhandled error type ({type(exc_repr)}), please report to"
" neotest-python repo"
)
result: NeotestResult = self.adapter.update_result( result: NeotestResult = self.adapter.update_result(
self.results.get(pos_id), self.results.get(pos_id),
{ {
@@ -199,7 +230,6 @@ class NeotestDebugpyPlugin:
# Do nothing if not running with a DAP debugger, # Do nothing if not running with a DAP debugger,
# e.g. neotest was invoked with {strategy = dap} # e.g. neotest was invoked with {strategy = dap}
return return
thread = threading.current_thread() thread = threading.current_thread()
additional_info = py_db.set_additional_thread_info(thread) additional_info = py_db.set_additional_thread_info(thread)
additional_info.is_tracing += 1 additional_info.is_tracing += 1
+4 -4
View File
@@ -45,7 +45,7 @@ class UnittestNeotestAdapter(NeotestAdapter):
return [*args, ".".join([relative_dotted, *child_ids])] return [*args, ".".join([relative_dotted, *child_ids])]
# TODO: Stream results # TODO: Stream results
def run(self, args: List[str], _) -> Dict: def run(self, args: List[str], _) -> Tuple[Dict, int]:
results = {} results = {}
errs: Dict[str, Tuple[Exception, Any, TracebackType]] = {} errs: Dict[str, Tuple[Exception, Any, TracebackType]] = {}
@@ -96,11 +96,11 @@ class UnittestNeotestAdapter(NeotestAdapter):
# Prepend an executable name which is just used in output # Prepend an executable name which is just used in output
argv = ["neotest-python"] + self.convert_args(args[-1], args[:-1]) argv = ["neotest-python"] + self.convert_args(args[-1], args[:-1])
unittest.main( program = unittest.main(
module=None, module=None,
argv=argv, argv=argv,
testRunner=NeotestUnittestRunner(resultclass=NeotestTextTestResult), testRunner=NeotestUnittestRunner(resultclass=NeotestTextTestResult),
exit=False, exit=False,
) )
exit_code = 0 if program.result.wasSuccessful() else 1
return results return results, exit_code