diff --git a/python/src/com/jetbrains/python/codeInsight/liveTemplates/PythonTemplateContextType.java b/python/src/com/jetbrains/python/codeInsight/liveTemplates/PythonTemplateContextType.java
index 11734409377b..23d21a6bbf59 100644
--- a/python/src/com/jetbrains/python/codeInsight/liveTemplates/PythonTemplateContextType.java
+++ b/python/src/com/jetbrains/python/codeInsight/liveTemplates/PythonTemplateContextType.java
@@ -2,6 +2,7 @@
package com.jetbrains.python.codeInsight.liveTemplates;
import com.intellij.codeInsight.template.EverywhereContextType;
+import com.intellij.codeInsight.template.TemplateActionContext;
import com.intellij.codeInsight.template.TemplateContextType;
import com.intellij.openapi.util.NlsContexts;
import com.intellij.patterns.PsiElementPattern;
@@ -31,19 +32,19 @@ public abstract class PythonTemplateContextType extends TemplateContextType {
}
@Override
- public boolean isInContext(@NotNull PsiFile file, int offset) {
- if (isPythonLanguage(file, offset)) {
- final PsiElement element = file.findElementAt(offset);
-
+ public boolean isInContext(@NotNull TemplateActionContext templateActionContext) {
+ PsiFile file = templateActionContext.getFile();
+ if (isPythonLanguage(file, templateActionContext.getStartOffset())) {
+ final PsiElement element = file.findElementAt(templateActionContext.getStartOffset());
if (element != null) {
- if (isAfterDot(element) || element instanceof PsiComment || isInsideStringLiteral(element) || isInsideParameterList(element)) {
- return false;
+ if (!templateActionContext.isSurrounding()) {
+ if (isAfterDot(element) || element instanceof PsiComment || isInsideStringLiteral(element) || isInsideParameterList(element)) {
+ return false;
+ }
}
-
return isInContext(element);
}
}
-
return false;
}
diff --git a/python/testData/codeInsight/liveTemplates/context/surroundStringLiteral.py b/python/testData/codeInsight/liveTemplates/context/surroundStringLiteral.py
new file mode 100644
index 000000000000..e54b04511f6b
--- /dev/null
+++ b/python/testData/codeInsight/liveTemplates/context/surroundStringLiteral.py
@@ -0,0 +1 @@
+'test'
\ No newline at end of file
diff --git a/python/testSrc/com/jetbrains/python/codeInsight/liveTemplates/PyLiveTemplatesContextTest.java b/python/testSrc/com/jetbrains/python/codeInsight/liveTemplates/PyLiveTemplatesContextTest.java
index 6e9d2a1bf953..280717fc5073 100644
--- a/python/testSrc/com/jetbrains/python/codeInsight/liveTemplates/PyLiveTemplatesContextTest.java
+++ b/python/testSrc/com/jetbrains/python/codeInsight/liveTemplates/PyLiveTemplatesContextTest.java
@@ -15,13 +15,14 @@
*/
package com.jetbrains.python.codeInsight.liveTemplates;
-import com.intellij.codeInsight.template.TemplateContextType;
+import com.intellij.codeInsight.template.*;
+import com.intellij.codeInsight.template.impl.*;
+import com.intellij.testFramework.fixtures.CodeInsightTestUtil;
+import com.intellij.util.containers.ContainerUtil;
import com.jetbrains.python.fixtures.PyTestCase;
import org.jetbrains.annotations.NotNull;
-import java.util.Arrays;
-import java.util.Comparator;
-import java.util.List;
+import java.util.*;
import java.util.stream.Collectors;
public class PyLiveTemplatesContextTest extends PyTestCase {
@@ -32,7 +33,7 @@ public class PyLiveTemplatesContextTest extends PyTestCase {
}
public void testNotPython() {
- doTest("html");
+ doTest(false, "html");
}
// PY-12212
@@ -64,22 +65,48 @@ public class PyLiveTemplatesContextTest extends PyTestCase {
doTest(PythonTemplateContextType.Class.class, PythonTemplateContextType.General.class);
}
- private void doTest(Class extends PythonTemplateContextType> @NotNull ... expectedContextTypes) {
- doTest("py", expectedContextTypes);
+ // PY-52162
+ public void testSurroundStringLiteral() {
+ doTest(true, "py", PythonTemplateContextType.General.class, PythonTemplateContextType.TopLevel.class);
}
- private void doTest(@NotNull String extension, Class extends PythonTemplateContextType> @NotNull ... expectedContextTypes) {
+ // PY-52162
+ public void testSurroundTemplateWithStringLiterals() {
+ final TemplateManager templateManager = TemplateManager.getInstance(myFixture.getProject());
+ final Template template = templateManager.createTemplate("pri", "Python", "print($SELECTION$)$END$");
+
+ TemplateContextType context = ContainerUtil.findInstance(TemplateContextType.EP_NAME.getExtensions(), PythonTemplateContextType.General.class);
+ assertNotNull(context);
+ ((TemplateImpl)template).getTemplateContext().setEnabled(context, true);
+
+ CodeInsightTestUtil.addTemplate(template, myFixture.getTestRootDisposable());
+
+ myFixture.configureByText("abc.py", "'test'");
+ myFixture.getEditor().getSelectionModel().setSelection(0, 6);
+ new InvokeTemplateAction((TemplateImpl)template, myFixture.getEditor(), myFixture.getProject(), new HashSet<>()).perform();
+ myFixture.checkResult("print('test')");
+ }
+
+ private void doTest(Class extends PythonTemplateContextType> @NotNull ... expectedContextTypes) {
+ doTest(false, "py", expectedContextTypes);
+ }
+
+ private void doTest(boolean isSurrounding,
+ @NotNull String extension,
+ Class extends PythonTemplateContextType> @NotNull ... expectedContextTypes) {
myFixture.configureByFile(getTestName(true) + "." + extension);
- final List> actualContextTypes = calculateEnabledContextTypes(getRegisteredContextTypes());
+ final List> actualContextTypes =
+ calculateEnabledContextTypes(isSurrounding, getRegisteredContextTypes());
assertSameElements(actualContextTypes, expectedContextTypes);
}
@NotNull
- private List> calculateEnabledContextTypes(@NotNull List registeredContextTypes) {
+ private List> calculateEnabledContextTypes(boolean isSurrounding, @NotNull List registeredContextTypes) {
+ TemplateActionContext context = TemplateActionContext.create(myFixture.getFile(), null, myFixture.getCaretOffset(), myFixture.getCaretOffset(), isSurrounding);
return registeredContextTypes
.stream()
- .filter(type -> type.isInContext(myFixture.getFile(), myFixture.getCaretOffset()))
+ .filter(type -> type.isInContext(context))
.map(PythonTemplateContextType::getClass)
.sorted(Comparator.comparing(Class::getSimpleName))
.collect(Collectors.toList());